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