Skip to content

Commit 405bc33

Browse files
committed
Merge remote-tracking branch 'template/main' into shimmy_boilerplate
2 parents 5edc022 + 8178698 commit 405bc33

2 files changed

Lines changed: 13 additions & 37 deletions

File tree

src/agent/llm_factory.py

Lines changed: 12 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1,36 +1,27 @@
11
import os
22
from typing import Optional
33

4-
from langchain_openai import AzureChatOpenAI
5-
from langchain_openai import AzureOpenAIEmbeddings
6-
from langchain_community.llms import Ollama
7-
from langchain_community.embeddings import OllamaEmbeddings
8-
from langchain_openai import ChatOpenAI
9-
from langchain_openai import OpenAIEmbeddings
10-
from langchain_google_genai import ChatGoogleGenerativeAI
114
from dotenv import load_dotenv
125
load_dotenv()
136

147
class AzureLLMs:
158
def __init__(self, temperature: int = 0):
9+
from langchain_openai import AzureChatOpenAI
10+
1611
self._azure_llm = AzureChatOpenAI(
1712
openai_api_version=os.environ["AZURE_OPENAI_API_VERSION"],
1813
azure_deployment=os.environ["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
1914
temperature=temperature,
2015
max_tokens=None,
2116
)
22-
self._azure_embedding = AzureOpenAIEmbeddings(azure_deployment=os.environ['AZURE_OPENAI_EMBEDDING_1536_DEPLOYMENT'],
23-
openai_api_version=os.environ["AZURE_OPENAI_API_VERSION"],
24-
model=os.environ["AZURE_OPENAI_EMBEDDING_1536_MODEL"])
25-
17+
2618
def get_llm(self):
2719
return self._azure_llm
2820

29-
def get_embedding(self):
30-
return self._azure_embedding
31-
3221
class OllamaLLMs:
3322
def __init__(self):
23+
from langchain_community.llms import Ollama
24+
3425
self._ollama_llm = Ollama(
3526
model=os.environ['OLLAMA_MODEL'],
3627
base_url=os.environ['OLLAMA_BASE_URL'],
@@ -39,42 +30,26 @@ def __init__(self):
3930
},
4031
)
4132

42-
self._ollama_embedding = OllamaEmbeddings(
43-
model='nomic-embed-text:137m-v1.5-fp16',
44-
base_url=os.environ['OLLAMA_BASE_URL'],
45-
headers={
46-
'X-API-Key': os.environ['OLLAMA_API_KEY'],
47-
},
48-
show_progress=True
49-
)
50-
5133
def get_llm(self):
5234
return self._ollama_llm
5335

54-
def get_embedding(self):
55-
return self._ollama_embedding
56-
5736
class OpenAILLMs:
5837
def __init__(self, temperature: int = 0):
38+
from langchain_openai import ChatOpenAI
39+
5940
self._openai_llm = ChatOpenAI(
6041
model=os.environ['OPENAI_MODEL'],
6142
temperature=temperature,
6243
api_key=os.environ["OPENAI_API_KEY"],
6344
)
6445

65-
self._openai_embedding = OpenAIEmbeddings(
66-
model='text-embedding-ada-002',
67-
api_key=os.environ['OPENAI_API_KEY'],
68-
)
69-
7046
def get_llm(self):
7147
return self._openai_llm
7248

73-
def get_embedding(self):
74-
return self._openai_embedding
75-
7649
class GoogleAILLMs:
7750
def __init__(self, temperature: int = 0):
51+
from langchain_google_genai import ChatGoogleGenerativeAI
52+
7853
self._google_llm = ChatGoogleGenerativeAI(
7954
model=os.environ['GOOGLE_AI_MODEL'],
8055
temperature=temperature,
@@ -84,8 +59,10 @@ def __init__(self, temperature: int = 0):
8459
def get_llm(self):
8560
return self._google_llm
8661

87-
class ChatOpenRouterProvider:
62+
class OpenRouterLLMs:
8863
def __init__(self, temperature: int = 0, model: Optional[str] = None):
64+
from langchain_openai import ChatOpenAI
65+
8966
model_name = model or os.environ['OPENROUTER_MODEL']
9067
key = os.environ['OPENROUTER_API_KEY']
9168
base_url = os.environ['OPENROUTER_BASE_URL']

src/module.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,6 @@
55
from lf_toolkit.shared.mued_api_v0_1_0 import DataPolicySupport, HealthStatus, Role
66

77
from src.agent.context import parse_json_to_prompt
8-
from src.agent.agent import invoke_base_agent
9-
108

119
def chat_module(request: ChatRequest) -> ChatResponse:
1210
"""
@@ -21,6 +19,7 @@ def chat_module(request: ChatRequest) -> ChatResponse:
2119
Edit src/agent/prompts.py to change the chatbot's behaviour.
2220
Edit src/agent/agent.py to change the agent logic (summarisation threshold, LLM provider, etc.).
2321
"""
22+
from src.agent.agent import invoke_base_agent
2423

2524
conversation_id = request.conversationId
2625

0 commit comments

Comments
 (0)