Bugfixes

Ollama context and chat model fix Endpoints in .env defer.py event loop fix Error logging for ui fix

frdel committed Sep 12, 2024 at 22:43 UTC 4307dce2c8f51bcbbb4228be0f6b2a91476fcc30
6 files changed +309 -197
agent.py
+257 -140
@@ -15,73 +15,79 @@ import python.helpers.log as Log
15 from python.helpers.dirty_json import DirtyJson
16 from python.helpers.defer import DeferredTask
17
18 +
19 class AgentContext:
20
20 - _contexts: dict[str, 'AgentContext'] = {}
21 + _contexts: dict[str, "AgentContext"] = {}
22 _counter: int = 0
22 -
23 - def __init__(self, config: 'AgentConfig', id:str|None = None, agent0: 'Agent|None' = None):
23 +
24 + def __init__(
25 + self, config: "AgentConfig", id: str | None = None, agent0: "Agent|None" = None
26 + ):
27 # build context
28 self.id = id or str(uuid.uuid4())
29 self.config = config
30 self.log = Log.Log()
31 self.agent0 = agent0 or Agent(0, self.config, self)
32 self.paused = False
30 - self.streaming_agent: Agent|None = None
31 - self.process: DeferredTask|None = None
33 + self.streaming_agent: Agent | None = None
34 + self.process: DeferredTask | None = None
35 AgentContext._counter += 1
33 - self.no = AgentContext._counter
36 + self.no = AgentContext._counter
37
38 self._contexts[self.id] = self
39
40 @staticmethod
38 - def get(id:str):
41 + def get(id: str):
42 return AgentContext._contexts.get(id, None)
43
44 @staticmethod
45 def first():
43 - if not AgentContext._contexts: return None
46 + if not AgentContext._contexts:
47 + return None
48 return list(AgentContext._contexts.values())[0]
49
46 -
50 @staticmethod
48 - def remove(id:str):
51 + def remove(id: str):
52 context = AgentContext._contexts.pop(id, None)
50 - if context and context.process: context.process.kill()
53 + if context and context.process:
54 + context.process.kill()
55 return context
56
57 def reset(self):
54 - if self.process: self.process.kill()
58 + if self.process:
59 + self.process.kill()
60 self.log.reset()
61 self.agent0 = Agent(0, self.config, self)
62 self.streaming_agent = None
58 - self.paused = False
63 + self.paused = False
64
60 -
65 def communicate(self, msg: str, broadcast_level: int = 1):
62 - self.paused=False #unpause if paused
63 -
66 + self.paused = False # unpause if paused
67 +
68 if self.process and self.process.is_alive():
65 - if self.streaming_agent: current_agent = self.streaming_agent
66 - else: current_agent = self.agent0
69 + if self.streaming_agent:
70 + current_agent = self.streaming_agent
71 + else:
72 + current_agent = self.agent0
73
74 # set intervention messages to agent(s):
75 intervention_agent = current_agent
70 - while intervention_agent and broadcast_level !=0:
76 + while intervention_agent and broadcast_level != 0:
77 intervention_agent.intervention_message = msg
78 broadcast_level -= 1
73 - intervention_agent = intervention_agent.data.get("superior",None)
79 + intervention_agent = intervention_agent.data.get("superior", None)
80 else:
81 self.process = DeferredTask(self.agent0.message_loop, msg)
82
83 return self.process
78 -
79 -
84 +
85 +
86 @dataclass
81 -class AgentConfig:
87 +class AgentConfig:
88 chat_model: BaseChatModel | BaseLLM
89 utility_model: BaseChatModel | BaseLLM
84 - embeddings_model:Embeddings
90 + embeddings_model: Embeddings
91 prompts_subdir: str = ""
92 memory_subdir: str = ""
93 knowledge_subdir: str = ""
@@ -99,8 +105,14 @@ class AgentConfig:
105 code_exec_docker_enabled: bool = True
106 code_exec_docker_name: str = "agent-zero-exe"
107 code_exec_docker_image: str = "frdel/agent-zero-exe:latest"
102 - code_exec_docker_ports: dict[str,int] = field(default_factory=lambda: {"22/tcp": 50022})
103 - code_exec_docker_volumes: dict[str, dict[str, str]] = field(default_factory=lambda: {files.get_abs_path("work_dir"): {"bind": "/root", "mode": "rw"}})
108 + code_exec_docker_ports: dict[str, int] = field(
109 + default_factory=lambda: {"22/tcp": 50022}
110 + )
111 + code_exec_docker_volumes: dict[str, dict[str, str]] = field(
112 + default_factory=lambda: {
113 + files.get_abs_path("work_dir"): {"bind": "/root", "mode": "rw"}
114 + }
115 + )
116 code_exec_ssh_enabled: bool = True
117 code_exec_ssh_addr: str = "localhost"
118 code_exec_ssh_port: int = 50022
@@ -108,20 +120,25 @@ class AgentConfig:
120 code_exec_ssh_pass: str = "toor"
121 additional: Dict[str, Any] = field(default_factory=dict)
122
123 +
124 # intervention exception class - skips rest of message loop iteration
125 class InterventionException(Exception):
126 pass
127
115 -# killer exception class - not forwarded to LLM, cannot be fixed on its own, ends message loop
116 -class KillerException(Exception):
128 +
129 +# repairable exception class - forwarded to LLM, may be fixed on its own
130 +class RepairableException(Exception):
131 pass
132
133 +
134 class Agent:
120 -
121 - def __init__(self, number:int, config: AgentConfig, context: AgentContext|None = None):
135
123 - # agent config
124 - self.config = config
136 + def __init__(
137 + self, number: int, config: AgentConfig, context: AgentContext | None = None
138 + ):
139 +
140 + # agent config
141 + self.config = config
142
143 # agent context
144 self.context = context or AgentContext(config)
@@ -133,104 +150,162 @@ class Agent:
150 self.history = []
151 self.last_message = ""
152 self.intervention_message = ""
136 - self.rate_limiter = rate_limiter.RateLimiter(self.context.log,max_calls=self.config.rate_limit_requests,max_input_tokens=self.config.rate_limit_input_tokens,max_output_tokens=self.config.rate_limit_output_tokens,window_seconds=self.config.rate_limit_seconds)
137 - self.data = {} # free data object all the tools can use
153 + self.rate_limiter = rate_limiter.RateLimiter(
154 + self.context.log,
155 + max_calls=self.config.rate_limit_requests,
156 + max_input_tokens=self.config.rate_limit_input_tokens,
157 + max_output_tokens=self.config.rate_limit_output_tokens,
158 + window_seconds=self.config.rate_limit_seconds,
159 + )
160 + self.data = {} # free data object all the tools can use
161
162 async def message_loop(self, msg: str):
163 try:
141 - printer = PrintStyle(italic=True, font_color="#b3ffd9", padding=False)
164 + printer = PrintStyle(italic=True, font_color="#b3ffd9", padding=False)
165 user_message = self.read_prompt("fw.user_message.md", message=msg)
143 - await self.append_message(user_message, human=True) # Append the user's input to the history
166 + await self.append_message(
167 + user_message, human=True
168 + ) # Append the user's input to the history
169 memories = await self.fetch_memories(True)
145 -
146 - while True: # let the agent iterate on his thoughts until he stops by using a tool
147 - self.context.streaming_agent = self #mark self as current streamer
170 +
171 + while (
172 + True
173 + ): # let the agent iterate on his thoughts until he stops by using a tool
174 + self.context.streaming_agent = self # mark self as current streamer
175 agent_response = ""
176
177 try:
178
152 - system = self.read_prompt("agent.system.md", agent_name=self.agent_name) + "\n\n" + self.read_prompt("agent.tools.md")
179 + system = (
180 + self.read_prompt("agent.system.md", agent_name=self.agent_name)
181 + + "\n\n"
182 + + self.read_prompt("agent.tools.md")
183 + )
184 memories = await self.fetch_memories()
154 - if memories: system+= "\n\n"+memories
185 + if memories:
186 + system += "\n\n" + memories
187 +
188 + prompt = ChatPromptTemplate.from_messages(
189 + [
190 + SystemMessage(content=system),
191 + MessagesPlaceholder(variable_name="messages"),
192 + ]
193 + )
194
156 - prompt = ChatPromptTemplate.from_messages([
157 - SystemMessage(content=system),
158 - MessagesPlaceholder(variable_name="messages") ])
159 -
195 inputs = {"messages": self.history}
196 chain = prompt | self.config.chat_model
197
198 formatted_inputs = prompt.format(messages=self.history)
164 - tokens = int(len(formatted_inputs)/4)
199 + tokens = int(len(formatted_inputs) / 4)
200 self.rate_limiter.limit_call_and_input(tokens)
166 -
201 +
202 # output that the agent is starting
168 - PrintStyle(bold=True, font_color="green", padding=True, background_color="white").print(f"{self.agent_name}: Generating:")
169 - log = self.context.log.log(type="agent", heading=f"{self.agent_name}: Generating:")
170 -
203 + PrintStyle(
204 + bold=True,
205 + font_color="green",
206 + padding=True,
207 + background_color="white",
208 + ).print(f"{self.agent_name}: Generating:")
209 + log = self.context.log.log(
210 + type="agent", heading=f"{self.agent_name}: Generating:"
211 + )
212 +
213 async for chunk in chain.astream(inputs):
172 - await self.handle_intervention(agent_response) # wait for intervention and handle it, if paused
214 + await self.handle_intervention(
215 + agent_response
216 + ) # wait for intervention and handle it, if paused
217 +
218 + if isinstance(chunk, str):
219 + content = chunk
220 + elif hasattr(chunk, "content"):
221 + content = str(chunk.content)
222 + else:
223 + content = str(chunk)
224
174 - if isinstance(chunk, str): content = chunk
175 - elif hasattr(chunk, "content"): content = str(chunk.content)
176 - else: content = str(chunk)
177 -
225 if content:
179 - printer.stream(content) # output the agent response stream
180 - agent_response += content # concatenate stream into the response
226 + printer.stream(content) # output the agent response stream
227 + agent_response += (
228 + content # concatenate stream into the response
229 + )
230 self.log_from_stream(agent_response, log)
231
183 - self.rate_limiter.set_output_tokens(int(len(agent_response)/4)) # rough estimation
184 -
232 + self.rate_limiter.set_output_tokens(
233 + int(len(agent_response) / 4)
234 + ) # rough estimation
235 +
236 await self.handle_intervention(agent_response)
237
187 - if self.last_message == agent_response: #if assistant_response is the same as last message in history, let him know
188 - await self.append_message(agent_response) # Append the assistant's response to the history
238 + if (
239 + self.last_message == agent_response
240 + ): # if assistant_response is the same as last message in history, let him know
241 + await self.append_message(
242 + agent_response
243 + ) # Append the assistant's response to the history
244 warning_msg = self.read_prompt("fw.msg_repeat.md")
190 - await self.append_message(warning_msg, human=True) # Append warning message to the history
245 + await self.append_message(
246 + warning_msg, human=True
247 + ) # Append warning message to the history
248 PrintStyle(font_color="orange", padding=True).print(warning_msg)
249 self.context.log.log(type="warning", content=warning_msg)
250
194 - else: #otherwise proceed with tool
195 - await self.append_message(agent_response) # Append the assistant's response to the history
196 - tools_result = await self.process_tools(agent_response) # process tools requested in agent message
197 - if tools_result: #final response of message loop available
198 - return tools_result #break the execution if the task is done
251 + else: # otherwise proceed with tool
252 + await self.append_message(
253 + agent_response
254 + ) # Append the assistant's response to the history
255 + tools_result = await self.process_tools(
256 + agent_response
257 + ) # process tools requested in agent message
258 + if tools_result: # final response of message loop available
259 + return (
260 + tools_result # break the execution if the task is done
261 + )
262
263 except InterventionException as e:
201 - pass # intervention message has been handled in handle_intervention(), proceed with conversation loop
264 + pass # intervention message has been handled in handle_intervention(), proceed with conversation loop
265 except asyncio.CancelledError as e:
203 - PrintStyle(font_color="white", background_color="red", padding=True).print(f"Context {self.context.id} terminated during message loop")
204 - raise e # process cancelled from outside, kill the loop
205 - except KillerException as e:
206 - error_message = errors.format_error(e)
207 - self.context.log.log(type="error", content=error_message)
208 - raise e # kill the loop
209 - except Exception as e: # Forward other errors to the LLM, maybe it can fix them
266 + PrintStyle(
267 + font_color="white", background_color="red", padding=True
268 + ).print(f"Context {self.context.id} terminated during message loop")
269 + raise e # process cancelled from outside, kill the loop
270 + except RepairableException as e: # Forward repairable errors to the LLM, maybe it can fix them
271 error_message = errors.format_error(e)
211 - msg_response = self.read_prompt("fw.error.md", error=error_message) # error message template
272 + msg_response = self.read_prompt(
273 + "fw.error.md", error=error_message
274 + ) # error message template
275 await self.append_message(msg_response, human=True)
276 PrintStyle(font_color="red", padding=True).print(msg_response)
277 self.context.log.log(type="error", content=msg_response)
215 -
278 + except Exception as e: # Other exception kill the loop
279 + error_message = errors.format_error(e)
280 + PrintStyle(font_color="red", padding=True).print(error_message)
281 + self.context.log.log(type="error", content=error_message)
282 + raise e # kill the loop
283 +
284 finally:
217 - self.context.streaming_agent = None # unset current streamer
285 + self.context.streaming_agent = None # unset current streamer
286
219 - def read_prompt(self, file:str, **kwargs):
287 + def read_prompt(self, file: str, **kwargs):
288 content = ""
289 if self.config.prompts_subdir:
290 try:
223 - content = files.read_file(files.get_abs_path(f"./prompts/{self.config.prompts_subdir}/{file}"), **kwargs)
291 + content = files.read_file(
292 + files.get_abs_path(
293 + f"./prompts/{self.config.prompts_subdir}/{file}"
294 + ),
295 + **kwargs,
296 + )
297 except Exception as e:
298 pass
299 if not content:
227 - content = files.read_file(files.get_abs_path(f"./prompts/default/{file}"), **kwargs)
300 + content = files.read_file(
301 + files.get_abs_path(f"./prompts/default/{file}"), **kwargs
302 + )
303 return content
304
230 - def get_data(self, field:str):
305 + def get_data(self, field: str):
306 return self.data.get(field, None)
307
233 - def set_data(self, field:str, value):
308 + def set_data(self, field: str, value):
309 self.data[field] = value
310
311 async def append_message(self, msg: str, human: bool = False):
@@ -240,17 +315,21 @@ class Agent:
315 else:
316 new_message = HumanMessage(content=msg) if human else AIMessage(content=msg)
317 self.history.append(new_message)
243 - await self.cleanup_history(self.config.msgs_keep_max, self.config.msgs_keep_start, self.config.msgs_keep_end)
244 - if message_type=="ai":
318 + await self.cleanup_history(
319 + self.config.msgs_keep_max,
320 + self.config.msgs_keep_start,
321 + self.config.msgs_keep_end,
322 + )
323 + if message_type == "ai":
324 self.last_message = msg
325
247 - def concat_messages(self,messages):
326 + def concat_messages(self, messages):
327 return "\n".join([f"{msg.type}: {msg.content}" for msg in messages])
328
250 - async def send_adhoc_message(self, system: str, msg: str, output_label:str):
251 - prompt = ChatPromptTemplate.from_messages([
252 - SystemMessage(content=system),
253 - HumanMessage(content=msg)])
329 + async def send_adhoc_message(self, system: str, msg: str, output_label: str):
330 + prompt = ChatPromptTemplate.from_messages(
331 + [SystemMessage(content=system), HumanMessage(content=msg)]
332 + )
333
334 chain = prompt | self.config.utility_model
335 response = ""
@@ -258,40 +337,54 @@ class Agent:
337 logger = None
338
339 if output_label:
261 - PrintStyle(bold=True, font_color="orange", padding=True, background_color="white").print(f"{self.agent_name}: {output_label}:")
340 + PrintStyle(
341 + bold=True, font_color="orange", padding=True, background_color="white"
342 + ).print(f"{self.agent_name}: {output_label}:")
343 printer = PrintStyle(italic=True, font_color="orange", padding=False)
263 - logger = self.context.log.log(type="adhoc", heading=f"{self.agent_name}: {output_label}:")
344 + logger = self.context.log.log(
345 + type="adhoc", heading=f"{self.agent_name}: {output_label}:"
346 + )
347
348 formatted_inputs = prompt.format()
266 - tokens = int(len(formatted_inputs)/4)
349 + tokens = int(len(formatted_inputs) / 4)
350 self.rate_limiter.limit_call_and_input(tokens)
268 -
351 +
352 async for chunk in chain.astream({}):
270 - if self.handle_intervention(): break # wait for intervention and handle it, if paused
353 + if self.handle_intervention():
354 + break # wait for intervention and handle it, if paused
355
272 - if isinstance(chunk, str): content = chunk
273 - elif hasattr(chunk, "content"): content = str(chunk.content)
274 - else: content = str(chunk)
356 + if isinstance(chunk, str):
357 + content = chunk
358 + elif hasattr(chunk, "content"):
359 + content = str(chunk.content)
360 + else:
361 + content = str(chunk)
362
276 - if printer: printer.stream(content)
277 - response+=content
278 - if logger: logger.update(content=response)
363 + if printer:
364 + printer.stream(content)
365 + response += content
366 + if logger:
367 + logger.update(content=response)
368
280 - self.rate_limiter.set_output_tokens(int(len(response)/4))
369 + self.rate_limiter.set_output_tokens(int(len(response) / 4))
370
371 return response
283 -
372 +
373 def get_last_message(self):
374 if self.history:
375 return self.history[-1]
376
288 - async def replace_middle_messages(self,middle_messages):
377 + async def replace_middle_messages(self, middle_messages):
378 cleanup_prompt = self.read_prompt("fw.msg_cleanup.md")
290 - summary = await self.send_adhoc_message(system=cleanup_prompt,msg=self.concat_messages(middle_messages), output_label="Mid messages cleanup summary")
379 + summary = await self.send_adhoc_message(
380 + system=cleanup_prompt,
381 + msg=self.concat_messages(middle_messages),
382 + output_label="Mid messages cleanup summary",
383 + )
384 new_human_message = HumanMessage(content=summary)
385 return [new_human_message]
386
294 - async def cleanup_history(self, max:int, keep_start:int, keep_end:int):
387 + async def cleanup_history(self, max: int, keep_start: int, keep_end: int):
388 if len(self.history) <= max:
389 return self.history
390
@@ -317,14 +410,24 @@ class Agent:
410
411 return self.history
412
320 - async def handle_intervention(self, progress:str=""):
321 - while self.context.paused: await asyncio.sleep(0.1) # wait if paused
322 - if self.intervention_message: # if there is an intervention message, but not yet processed
413 + async def handle_intervention(self, progress: str = ""):
414 + while self.context.paused:
415 + await asyncio.sleep(0.1) # wait if paused
416 + if (
417 + self.intervention_message
418 + ): # if there is an intervention message, but not yet processed
419 msg = self.intervention_message
324 - self.intervention_message = "" # reset the intervention message
325 - if progress.strip(): await self.append_message(progress) # append the response generated so far
326 - user_msg = self.read_prompt("fw.intervention.md", user_message=msg) # format the user intervention template
327 - await self.append_message(user_msg,human=True) # append the intervention message
420 + self.intervention_message = "" # reset the intervention message
421 + if progress.strip():
422 + await self.append_message(
423 + progress
424 + ) # append the response generated so far
425 + user_msg = self.read_prompt(
426 + "fw.intervention.md", user_message=msg
427 + ) # format the user intervention template
428 + await self.append_message(
429 + user_msg, human=True
430 + ) # append the intervention message
431 raise InterventionException(msg)
432
433 async def process_tools(self, msg: str):
@@ -335,30 +438,36 @@ class Agent:
438 tool_name = tool_request.get("tool_name", "")
439 tool_args = tool_request.get("tool_args", {})
440 tool = self.get_tool(tool_name, tool_args, msg)
338 -
339 - await self.handle_intervention() # wait if paused and handle intervention message if needed
441 +
442 + await self.handle_intervention() # wait if paused and handle intervention message if needed
443 await tool.before_execution(**tool_args)
341 - await self.handle_intervention() # wait if paused and handle intervention message if needed
444 + await self.handle_intervention() # wait if paused and handle intervention message if needed
445 response = await tool.execute(**tool_args)
343 - await self.handle_intervention() # wait if paused and handle intervention message if needed
446 + await self.handle_intervention() # wait if paused and handle intervention message if needed
447 await tool.after_execution(response)
345 - await self.handle_intervention() # wait if paused and handle intervention message if needed
346 - if response.break_loop: return response.message
448 + await self.handle_intervention() # wait if paused and handle intervention message if needed
449 + if response.break_loop:
450 + return response.message
451 else:
452 msg = self.read_prompt("fw.msg_misformat.md")
453 await self.append_message(msg, human=True)
454 PrintStyle(font_color="red", padding=True).print(msg)
351 - self.context.log.log(type="error", content=f"{self.agent_name}: Message misformat:")
352 -
455 + self.context.log.log(
456 + type="error", content=f"{self.agent_name}: Message misformat:"
457 + )
458
459 def get_tool(self, name: str, args: dict, message: str, **kwargs):
355 - from python.tools.unknown import Unknown
460 + from python.tools.unknown import Unknown
461 from python.helpers.tool import Tool
357 -
462 +
463 tool_class = Unknown
359 - if files.exists("python/tools",f"{name}.py"):
360 - module = importlib.import_module("python.tools." + name) # Import the module
361 - class_list = inspect.getmembers(module, inspect.isclass) # Get all functions in the module
464 + if files.exists("python/tools", f"{name}.py"):
465 + module = importlib.import_module(
466 + "python.tools." + name
467 + ) # Import the module
468 + class_list = inspect.getmembers(
469 + module, inspect.isclass
470 + ) # Get all functions in the module
471
472 for cls in class_list:
473 if cls[1] is not Tool and issubclass(cls[1], Tool):
@@ -367,33 +476,41 @@ class Agent:
476
477 return tool_class(agent=self, name=name, args=args, message=message, **kwargs)
478
370 - async def fetch_memories(self,reset_skip=False):
371 - if self.config.auto_memory_count<=0: return ""
372 - if reset_skip: self.memory_skip_counter = 0
479 + async def fetch_memories(self, reset_skip=False):
480 + if self.config.auto_memory_count <= 0:
481 + return ""
482 + if reset_skip:
483 + self.memory_skip_counter = 0
484
485 if self.memory_skip_counter > 0:
375 - self.memory_skip_counter-=1
486 + self.memory_skip_counter -= 1
487 return ""
488 else:
489 self.memory_skip_counter = self.config.auto_memory_skip
490 from python.tools import memory_tool
491 +
492 messages = self.concat_messages(self.history)
381 - memories = memory_tool.search(self,messages)
382 - input = {
383 - "conversation_history" : messages,
384 - "raw_memories": memories
385 - }
386 - cleanup_prompt = self.read_prompt("msg.memory_cleanup.md").replace("{", "{{")
387 - clean_memories = await self.send_adhoc_message(cleanup_prompt,json.dumps(input), output_label="Memory injection")
493 + memories = memory_tool.search(self, messages)
494 + input = {"conversation_history": messages, "raw_memories": memories}
495 + cleanup_prompt = self.read_prompt("msg.memory_cleanup.md").replace(
496 + "{", "{{"
497 + )
498 + clean_memories = await self.send_adhoc_message(
499 + cleanup_prompt, json.dumps(input), output_label="Memory injection"
500 + )
501 return clean_memories
502
503 def log_from_stream(self, stream: str, logItem: Log.LogItem):
504 try:
392 - if len(stream) < 25: return # no reason to try
505 + if len(stream) < 25:
506 + return # no reason to try
507 response = DirtyJson.parse_string(stream)
394 - if isinstance(response, dict): logItem.update(content=stream, kvps=response) #log if result is a dictionary already
508 + if isinstance(response, dict):
509 + logItem.update(
510 + content=stream, kvps=response
511 + ) # log if result is a dictionary already
512 except Exception as e:
513 pass
514
515 def call_extension(self, name: str, **kwargs) -> Any:
399 - pass
\ No newline at end of file
516 + pass
example.env
+6 -2
@@ -1,4 +1,4 @@
1 -API_KEY_OPENAI=
1 +API_KEY_OPENAI=sk-hyBlbkFJCJjaYGCbqPTyT3uaYGCbqFBlbkFJCyJCyuPhYGCb
2 API_KEY_ANTHROPIC=
3 API_KEY_GROQ=
4 API_KEY_PERPLEXITY=
@@ -16,4 +16,8 @@ WEB_UI_PORT=50001
16
17
18 TOKENIZERS_PARALLELISM=true
19 -PYDEVD_DISABLE_FILE_VALIDATION=1
\ No newline at end of file
19 +PYDEVD_DISABLE_FILE_VALIDATION=1
20 +
21 +OLLAMA_BASE_URL="http://127.0.0.1:11434"
22 +LM_STUDIO_BASE_URL="http://127.0.0.1:1234/v1"
23 +OPEN_ROUTER_BASE_URL="https://openrouter.ai/api/v1"
\ No newline at end of file
initialize.py
+1 -1
@@ -7,7 +7,7 @@ def initialize():
7 chat_llm = models.get_openai_chat(model_name="gpt-4o-mini", temperature=0)
8 # chat_llm = models.get_ollama_chat(model_name="gemma2:latest", temperature=0)
9 # chat_llm = models.get_lmstudio_chat(model_name="TheBloke/Mistral-7B-Instruct-v0.2-GGUF", temperature=0)
10 - # chat_llm = models.get_openrouter(model_name="meta-llama/llama-3-8b-instruct:free")
10 + # chat_llm = models.get_openrouter_chat(model_name="nousresearch/hermes-3-llama-3.1-405b")
11 # chat_llm = models.get_azure_openai_chat(deployment_name="gpt-4o-mini", temperature=0)
12 # chat_llm = models.get_anthropic_chat(model_name="claude-3-5-sonnet-20240620", temperature=0)
13 # chat_llm = models.get_google_chat(model_name="gemini-1.5-flash", temperature=0)
models.py
+24 -39
@@ -2,11 +2,12 @@ import os
2 from dotenv import load_dotenv
3 from langchain_openai import ChatOpenAI, OpenAI, OpenAIEmbeddings, AzureChatOpenAI, AzureOpenAIEmbeddings, AzureOpenAI
4 from langchain_community.llms.ollama import Ollama
5 +from langchain_ollama import ChatOllama
6 from langchain_community.embeddings import OllamaEmbeddings
7 from langchain_anthropic import ChatAnthropic
8 from langchain_groq import ChatGroq
9 from langchain_huggingface import HuggingFaceEmbeddings
9 -from langchain_google_genai import ChatGoogleGenerativeAI, HarmBlockThreshold, HarmCategory
10 +from langchain_google_genai import GoogleGenerativeAI, HarmBlockThreshold, HarmCategory
11 from pydantic.v1.types import SecretStr
12
13
@@ -22,11 +23,12 @@ def get_api_key(service):
23
24
25 # Ollama models
25 -def get_ollama_chat(model_name:str, temperature=DEFAULT_TEMPERATURE, base_url="http://localhost:11434"):
26 - return Ollama(model=model_name,temperature=temperature, base_url=base_url)
26 +def get_ollama_chat(model_name:str, temperature=DEFAULT_TEMPERATURE, base_url=os.getenv("OLLAMA_BASE_URL") or "http://127.0.0.1:11434", num_ctx=8192):
27 + return ChatOllama(model=model_name,temperature=temperature, base_url=base_url, num_ctx=num_ctx)
28
28 -def get_ollama_embedding(model_name:str, temperature=DEFAULT_TEMPERATURE):
29 - return OllamaEmbeddings(model=model_name,temperature=temperature)
29 +def get_ollama_embedding(model_name:str, temperature=DEFAULT_TEMPERATURE, base_url=os.getenv("OLLAMA_BASE_URL") or "http://127.0.0.1:11434"):
30 +
31 + return OllamaEmbeddings(model=model_name,temperature=temperature, base_url=base_url)
32
33 # HuggingFace models
34
@@ -34,63 +36,46 @@ def get_huggingface_embedding(model_name:str):
36 return HuggingFaceEmbeddings(model_name=model_name)
37
38 # LM Studio and other OpenAI compatible interfaces
37 -def get_lmstudio_chat(model_name:str, base_url="http://localhost:1234/v1", temperature=DEFAULT_TEMPERATURE):
39 +def get_lmstudio_chat(model_name:str, temperature=DEFAULT_TEMPERATURE, base_url=os.getenv("LM_STUDIO_BASE_URL") or "http://127.0.0.1:1234/v1"):
40 return ChatOpenAI(model_name=model_name, base_url=base_url, temperature=temperature, api_key="none") # type: ignore
41
40 -def get_lmstudio_embedding(model_name:str, base_url="http://localhost:1234/v1"):
42 +def get_lmstudio_embedding(model_name:str, base_url=os.getenv("LM_STUDIO_BASE_URL") or "http://127.0.0.1:1234/v1"):
43 return OpenAIEmbeddings(model_name=model_name, base_url=base_url) # type: ignore
44
45 # Anthropic models
44 -def get_anthropic_chat(model_name:str, api_key=None, temperature=DEFAULT_TEMPERATURE):
45 - api_key = api_key or get_api_key("anthropic")
46 +def get_anthropic_chat(model_name:str, api_key=get_api_key("anthropic"), temperature=DEFAULT_TEMPERATURE):
47 return ChatAnthropic(model_name=model_name, temperature=temperature, api_key=api_key) # type: ignore
48
49 # OpenAI models
49 -def get_openai_chat(model_name:str, api_key=None, temperature=DEFAULT_TEMPERATURE):
50 - api_key = api_key or get_api_key("openai")
50 +def get_openai_chat(model_name:str, api_key=get_api_key("openai"), temperature=DEFAULT_TEMPERATURE):
51 return ChatOpenAI(model_name=model_name, temperature=temperature, api_key=api_key) # type: ignore
52
53 -def get_openai_instruct(model_name:str,api_key=None, temperature=DEFAULT_TEMPERATURE):
54 - api_key = api_key or get_api_key("openai")
53 +def get_openai_instruct(model_name:str, api_key=get_api_key("openai"), temperature=DEFAULT_TEMPERATURE):
54 return OpenAI(model=model_name, temperature=temperature, api_key=api_key) # type: ignore
55
57 -def get_openai_embedding(model_name:str, api_key=None):
58 - api_key = api_key or get_api_key("openai")
56 +def get_openai_embedding(model_name:str, api_key=get_api_key("openai")):
57 return OpenAIEmbeddings(model=model_name, api_key=api_key) # type: ignore
58
61 -def get_azure_openai_chat(deployment_name:str, api_key=None, temperature=DEFAULT_TEMPERATURE, azure_endpoint=None):
62 - api_key = api_key or get_api_key("openai_azure")
63 - azure_endpoint = azure_endpoint or os.getenv("OPENAI_AZURE_ENDPOINT")
59 +def get_azure_openai_chat(deployment_name:str, api_key=get_api_key("openai_azure"), temperature=DEFAULT_TEMPERATURE, azure_endpoint=os.getenv("OPENAI_AZURE_ENDPOINT")):
60 return AzureChatOpenAI(deployment_name=deployment_name, temperature=temperature, api_key=api_key, azure_endpoint=azure_endpoint) # type: ignore
61
66 -def get_azure_openai_instruct(deployment_name:str, api_key=None, temperature=DEFAULT_TEMPERATURE, azure_endpoint=None):
67 - api_key = api_key or get_api_key("openai_azure")
68 - azure_endpoint = azure_endpoint or os.getenv("OPENAI_AZURE_ENDPOINT")
62 +def get_azure_openai_instruct(deployment_name:str, api_key=get_api_key("openai_azure"), temperature=DEFAULT_TEMPERATURE, azure_endpoint=os.getenv("OPENAI_AZURE_ENDPOINT")):
63 return AzureOpenAI(deployment_name=deployment_name, temperature=temperature, api_key=api_key, azure_endpoint=azure_endpoint) # type: ignore
64
71 -def get_azure_openai_embedding(deployment_name:str, api_key=None, azure_endpoint=None):
72 - api_key = api_key or get_api_key("openai_azure")
73 - azure_endpoint = azure_endpoint or os.getenv("OPENAI_AZURE_ENDPOINT")
65 +def get_azure_openai_embedding(deployment_name:str, api_key=get_api_key("openai_azure"), azure_endpoint=os.getenv("OPENAI_AZURE_ENDPOINT")):
66 return AzureOpenAIEmbeddings(deployment_name=deployment_name, api_key=api_key, azure_endpoint=azure_endpoint) # type: ignore
67
68 # Google models
77 -def get_google_chat(model_name:str, api_key=None, temperature=DEFAULT_TEMPERATURE):
78 - api_key = api_key or get_api_key("google")
79 - return ChatGoogleGenerativeAI(model=model_name, temperature=temperature, google_api_key=api_key, safety_settings={HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT: HarmBlockThreshold.BLOCK_NONE }) # type: ignore
69 +def get_google_chat(model_name:str, api_key=get_api_key("google"), temperature=DEFAULT_TEMPERATURE):
70 + return GoogleGenerativeAI(model=model_name, temperature=temperature, google_api_key=api_key, safety_settings={HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT: HarmBlockThreshold.BLOCK_NONE }) # type: ignore
71
72 # Groq models
82 -def get_groq_chat(model_name:str, api_key=None, temperature=DEFAULT_TEMPERATURE):
83 - api_key = api_key or get_api_key("groq")
73 +def get_groq_chat(model_name:str, api_key=get_api_key("groq"), temperature=DEFAULT_TEMPERATURE):
74 return ChatGroq(model_name=model_name, temperature=temperature, api_key=api_key) # type: ignore
75
76 # OpenRouter models
87 -def get_openrouter(model_name: str="meta-llama/llama-3.1-8b-instruct:free", api_key=None, temperature=DEFAULT_TEMPERATURE):
88 - api_key = api_key or get_api_key("openrouter")
89 - return ChatOpenAI(api_key=api_key, base_url="https://openrouter.ai/api/v1", model=model_name, temperature=temperature) # type: ignore
90 -
91 -def get_embedding_hf(model_name="sentence-transformers/all-MiniLM-L6-v2"):
92 - return HuggingFaceEmbeddings(model_name=model_name)
93 -
94 -def get_embedding_openai(api_key=None):
95 - api_key = api_key or get_api_key("openai")
96 - return OpenAIEmbeddings(api_key=api_key) #type: ignore
77 +def get_openrouter_chat(model_name: str, api_key=get_api_key("openrouter"), temperature=DEFAULT_TEMPERATURE, base_url=os.getenv("OPEN_ROUTER_BASE_URL") or "https://openrouter.ai/api/v1"):
78 + return ChatOpenAI(api_key=api_key, model=model_name, temperature=temperature, base_url=base_url) # type: ignore
79 +
80 +def get_openrouter_embedding(model_name: str, api_key=get_api_key("openrouter"), base_url=os.getenv("OPEN_ROUTER_BASE_URL") or "https://openrouter.ai/api/v1"):
81 + return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url) # type: ignore
\ No newline at end of file
python/helpers/defer.py
+20 -15
@@ -4,20 +4,21 @@ from concurrent.futures import Future
4
5 class DeferredTask:
6 def __init__(self, func, *args, **kwargs):
7 - self._loop = asyncio.new_event_loop()
7 + self._loop: asyncio.AbstractEventLoop = None # type: ignore
8 self._task = None
9 self._future = Future()
10 - self._task_initialized = threading.Event() # Event to signal task initialization
10 + self._task_initialized = threading.Event()
11 self._start_task(func, *args, **kwargs)
12
13 def _start_task(self, func, *args, **kwargs):
14 - def run_in_thread(loop, func, args, kwargs):
15 - asyncio.set_event_loop(loop)
16 - self._task = loop.create_task(self._run(func, *args, **kwargs))
17 - self._task_initialized.set() # Signal that the task has been initialized
18 - loop.run_forever()
14 + def run_in_thread():
15 + self._loop = asyncio.new_event_loop()
16 + asyncio.set_event_loop(self._loop)
17 + self._task = self._loop.create_task(self._run(func, *args, **kwargs))
18 + self._task_initialized.set()
19 + self._loop.run_forever()
20
20 - self._thread = threading.Thread(target=run_in_thread, args=(self._loop, func, args, kwargs))
21 + self._thread = threading.Thread(target=run_in_thread)
22 self._thread.start()
23
24 async def _run(self, func, *args, **kwargs):
@@ -27,13 +28,16 @@ class DeferredTask:
28 except Exception as e:
29 self._future.set_exception(e)
30 finally:
30 - self._loop.call_soon_threadsafe(self._loop.stop)
31 + self._loop.call_soon_threadsafe(self._cleanup)
32 +
33 + def _cleanup(self):
34 + self._loop.stop()
35
36 def is_ready(self):
37 return self._future.done()
38
39 async def result(self, timeout=None):
36 - if not self._task_initialized.wait(timeout): # Wait until the task is initialized
40 + if not self._task_initialized.wait(timeout):
41 raise RuntimeError("Task was not initialized properly.")
42
43 try:
@@ -42,7 +46,7 @@ class DeferredTask:
46 raise TimeoutError("The task did not complete within the specified timeout.")
47
48 def result_sync(self, timeout=None):
45 - if not self._task_initialized.wait(timeout): # Wait until the task is initialized
49 + if not self._task_initialized.wait(timeout):
50 raise RuntimeError("Task was not initialized properly.")
51
52 try:
@@ -58,8 +62,9 @@ class DeferredTask:
62 return self._thread.is_alive() and not self._future.done()
63
64 def __del__(self):
61 - if self._loop.is_running():
62 - self._loop.call_soon_threadsafe(self._loop.stop)
63 - if self._thread.is_alive():
65 + if self._loop and self._loop.is_running():
66 + self._loop.call_soon_threadsafe(self._cleanup)
67 + if self._thread and self._thread.is_alive():
68 self._thread.join()
65 - self._loop.close()
\ No newline at end of file
69 + if self._loop:
70 + self._loop.close()
\ No newline at end of file
requirements.txt
+1
@@ -2,6 +2,7 @@ ansio==0.0.1
2 python-dotenv==1.0.1
3 langchain-groq==0.1.6
4 langchain-huggingface==0.0.3
5 +langchain-ollama==0.1.3
6 langchain-openai==0.1.15
7 langchain-community==0.2.7
8 langchain-anthropic==0.1.19