llm-interface 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,427 @@
1
+ # Copyright 2024 Niels Provos
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ import hashlib
15
+ import json
16
+ import logging
17
+ import re
18
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Type
19
+
20
+ import diskcache
21
+ from dotenv import load_dotenv
22
+ from ollama import Client
23
+ from pydantic import BaseModel
24
+
25
+ from .llm_tool import Tool
26
+ from .pydantic_output_parser import MinimalPydanticOutputParser
27
+ from .utils import setup_logging
28
+
29
+ # Load environment variables from .env.local file
30
+ load_dotenv(".env.local")
31
+
32
+
33
+ class NoCache:
34
+ def get(self, key):
35
+ return None
36
+
37
+ def set(self, key, value):
38
+ pass
39
+
40
+
41
+ class LLMInterface:
42
+ def __init__(
43
+ self,
44
+ model_name: str = "llama2",
45
+ log_dir: str = "logs",
46
+ client: Optional[Any] = None,
47
+ host: Optional[str] = None,
48
+ support_json_mode: bool = True,
49
+ support_structured_outputs: bool = False,
50
+ support_system_prompt: bool = True,
51
+ use_cache: bool = True,
52
+ ):
53
+ self.model_name = model_name
54
+ self.client = client if client else Client(host=host)
55
+ self.support_json_mode = support_json_mode
56
+ self.support_structured_outputs = support_structured_outputs
57
+ self.support_system_prompt = support_system_prompt
58
+
59
+ self.logger = setup_logging(
60
+ logs_dir=log_dir, logs_prefix="llm_interface", logger_name=__name__
61
+ )
62
+
63
+ # Initialize disk cache for caching responses
64
+ self.disk_cache = (
65
+ diskcache.Cache(
66
+ directory=".response_cache", eviction_policy="least-recently-used"
67
+ )
68
+ if use_cache
69
+ else NoCache()
70
+ )
71
+
72
+ def _execute_tool(
73
+ self, tool_call: Dict[str, Any], tools: List[Tool]
74
+ ) -> List[Dict[str, str]]:
75
+ """Execute tool calls and format results for the conversation."""
76
+ tool_name = tool_call.get("name") or tool_call.get("function", {}).get("name")
77
+ arguments = tool_call.get("arguments") or tool_call.get("function", {}).get(
78
+ "arguments", {}
79
+ )
80
+
81
+ if isinstance(arguments, str):
82
+ # Parse JSON string if needed
83
+ try:
84
+ arguments = json.loads(arguments)
85
+ except json.JSONDecodeError:
86
+ self.logger.error("Failed to parse tool arguments: %s", arguments)
87
+ return []
88
+
89
+ tool_map = {tool.name: tool for tool in tools}
90
+ if tool_name in tool_map:
91
+ try:
92
+ result = tool_map[tool_name].execute(**arguments)
93
+ return [
94
+ {
95
+ "role": "assistant",
96
+ "content": "",
97
+ "tool_calls": [
98
+ {
99
+ "id": tool_call.get("id", ""),
100
+ "type": "function",
101
+ "function": {
102
+ "name": tool_name,
103
+ "arguments": arguments, # Ollama requires this as a Dict but OpenAI requires it as a string
104
+ },
105
+ }
106
+ ],
107
+ },
108
+ {
109
+ "role": "tool",
110
+ "name": tool_name,
111
+ "tool_call_id": tool_call.get("id", ""),
112
+ "content": str(result),
113
+ },
114
+ ]
115
+ except Exception as e:
116
+ self.logger.error("Tool execution failed: %s", e)
117
+ return []
118
+ else:
119
+ self.logger.error("Tool '%s' not found.", tool_name)
120
+ return []
121
+
122
+ def _cached_chat(
123
+ self,
124
+ messages: List[Dict[str, str]],
125
+ tools: Optional[List[Tool]] = None,
126
+ temperature: Optional[float] = None,
127
+ response_schema: Optional[Type[BaseModel]] = None,
128
+ ) -> str:
129
+ # Concatenate all messages to use as the cache key
130
+ message_content = "".join([msg["role"] + msg["content"] for msg in messages])
131
+ tool_content = ""
132
+ if tools:
133
+ tool_content = "".join([f"{t.name}{t.description}" for t in tools])
134
+ prompt_hash = self._generate_hash(
135
+ self.model_name + f"-{temperature}"
136
+ if temperature
137
+ else "" + message_content + tool_content
138
+ )
139
+
140
+ self.logger.info("Chatting with messages: %s", messages)
141
+
142
+ # Check if prompt response is in cache
143
+ response = self.disk_cache.get(prompt_hash)
144
+
145
+ if response is None:
146
+ kwargs = {}
147
+ current_messages = messages.copy()
148
+
149
+ # some models can generate structured outputs
150
+ if self.support_structured_outputs and response_schema:
151
+ if isinstance(self.client, Client):
152
+ # For Ollama, we need to pass the schema directly
153
+ kwargs["format"] = response_schema.model_json_schema()
154
+ else:
155
+ # For OpenAI, we need to pass the schema as a pydantic object
156
+ kwargs["response_schema"] = response_schema
157
+ elif self.support_json_mode:
158
+ kwargs["format"] = "json"
159
+
160
+ # ollama expects temperature to be passed as an option
161
+ options = {}
162
+ if temperature:
163
+ options["temperature"] = temperature
164
+ kwargs["options"] = options
165
+
166
+ num_tool_calls = 0
167
+ max_tool_calls = 5
168
+ while num_tool_calls < max_tool_calls:
169
+ num_tool_calls += 1
170
+
171
+ converted_tools = [tool.to_dict() for tool in tools] if tools else []
172
+ # Make request to client using chat interface
173
+ response = self.client.chat(
174
+ model=self.model_name,
175
+ tools=converted_tools,
176
+ messages=current_messages,
177
+ **kwargs,
178
+ )
179
+
180
+ self.logger.info("Received chat response: %s", response)
181
+
182
+ # Check if the response contains tool calls
183
+ tool_calls = response.get("message", {}).get("tool_calls", [])
184
+ if not tool_calls:
185
+ break
186
+
187
+ self.logger.info("Received tool calls: %s", tool_calls)
188
+ # Execute tools and add results to messages
189
+ for tool_call in tool_calls:
190
+ tool_messages = self._execute_tool(tool_call, tools)
191
+ current_messages.extend(tool_messages)
192
+
193
+ self.logger.info("Chatting with messages: %s", current_messages)
194
+
195
+ # Cache the response with hashed prompt as key
196
+ self.disk_cache.set(prompt_hash, response)
197
+
198
+ if "error" in response:
199
+ self.logger.error("Error in chat response: %s", response["error"])
200
+ elif "refusal" in response:
201
+ self.logger.error("Model refused the request: %s", response["refusal"])
202
+
203
+ return response["message"]["content"]
204
+
205
+ def chat(
206
+ self,
207
+ messages: List[Dict[str, str]],
208
+ tools: Optional[List[Tool]] = None,
209
+ temperature: Optional[float] = None,
210
+ response_schema: Optional[Type[BaseModel]] = None,
211
+ ) -> str:
212
+ response = self._cached_chat(
213
+ messages=messages,
214
+ tools=tools,
215
+ temperature=temperature,
216
+ response_schema=response_schema,
217
+ )
218
+ self.logger.info(
219
+ "Received chat response: %s...",
220
+ response[:850] if isinstance(response, str) else response,
221
+ )
222
+ return response.strip() if isinstance(response, str) else response
223
+
224
+ def _strip_text_from_json_response(self, response: str) -> str:
225
+ pattern = r"^[^{\[]*([{\[].*[}\]])[^}\]]*$"
226
+ match = re.search(pattern, response, re.DOTALL)
227
+
228
+ if match:
229
+ return match.group(1)
230
+ else:
231
+ return response # Return original response if no JSON block is found
232
+
233
+ def generate_full_prompt(
234
+ self, prompt_template: str, system: str = "", **kwargs
235
+ ) -> str:
236
+ """
237
+ Generate a full prompt with input variables filled in.
238
+
239
+ Args:
240
+ prompt_template (str): The prompt template with placeholders for variables.
241
+ system (str): The system prompt to use for generation.
242
+ **kwargs: Keyword arguments to fill in the prompt template.
243
+
244
+ Returns:
245
+ str: The formatted prompt
246
+ """
247
+ formatted_prompt = prompt_template.format(**kwargs)
248
+ return formatted_prompt
249
+
250
+ def generate_pydantic(
251
+ self,
252
+ prompt_template: str,
253
+ output_schema: Type[BaseModel],
254
+ system: str = "",
255
+ tools: Optional[List[Tool]] = None,
256
+ logger: Optional[logging.Logger] = None,
257
+ debug_saver: Optional[Callable[[str, Dict[str, Any], str], None]] = None,
258
+ extra_validation: Optional[Callable[[BaseModel], Optional[str]]] = None,
259
+ temperature: Optional[float] = None,
260
+ **kwargs,
261
+ ) -> Optional[BaseModel]:
262
+ """
263
+ Generates a Pydantic model instance based on a specified prompt template and output schema.
264
+
265
+ This function uses a prompt template with variable placeholders to generate a full prompt. It utilizes
266
+ this prompt in combination with a specified system prompt to interact with a chat-based interface,
267
+ aiming to produce a structured output conforming to a given Pydantic schema. The function attempts up to
268
+ three iterations to obtain a valid response, applying parsing, validation, and optional extra validation
269
+ functions. If all iterations fail, None is returned.
270
+
271
+ Args:
272
+ prompt_template (str): The template containing placeholders for formatting the prompt.
273
+ output_schema (Type[BaseModel]): A Pydantic model that defines the expected schema of the output data.
274
+ system (str): An optional system prompt used during the generation process.
275
+ logger (Optional[logging.Logger]): An optional logger for recording the generated prompt and events.
276
+ debug_saver (Optional[Callable[[str, Dict[str, Any], str], None]]): An optional callback for saving debugging information,
277
+ which receives the prompt and the response.
278
+ extra_validation (Optional[Callable[[BaseModel], str]]): An optional function for additional validation of
279
+ the generated output. It should return an error message if validation fails, otherwise None.
280
+ **kwargs: Additional keyword arguments for populating the prompt template.
281
+
282
+ Returns:
283
+ Optional[BaseModel]: An instance of the specified Pydantic model with generated data if successful,
284
+ or None if all attempts at generation fail or the response is invalid.
285
+ """
286
+ parser = MinimalPydanticOutputParser(pydantic_object=output_schema)
287
+
288
+ formatted_prompt = self.generate_full_prompt(
289
+ prompt_template=prompt_template, system=system, **kwargs
290
+ )
291
+
292
+ self.logger.info("Generated prompt: %s", formatted_prompt)
293
+ if logger:
294
+ logger.info("Generated prompt: %s", formatted_prompt)
295
+
296
+ messages = []
297
+ if self.support_system_prompt:
298
+ messages.append({"role": "system", "content": system})
299
+ messages.append({"role": "user", "content": formatted_prompt})
300
+
301
+ iteration = 0
302
+ while iteration < 3:
303
+ iteration += 1
304
+
305
+ raw_response = self.chat(
306
+ messages=messages,
307
+ temperature=temperature,
308
+ response_schema=output_schema,
309
+ tools=tools,
310
+ )
311
+
312
+ if self.support_structured_outputs:
313
+ # If the model supports structured outputs, we should get a Pydantic object directly
314
+ # or a string that can be parsed directly
315
+ try:
316
+ if raw_response is None:
317
+ response = None
318
+ error_message = "The model refused the request"
319
+ elif isinstance(raw_response, BaseModel):
320
+ response = raw_response
321
+ error_message = None
322
+ else:
323
+ response = output_schema.model_validate_json(raw_response)
324
+ error_message = None
325
+ except Exception as e:
326
+ self.logger.error("Error parsing structured response: %s", e)
327
+ error_message = str(e)
328
+ response = None
329
+ else:
330
+ if not self.support_json_mode:
331
+ raw_response = self._strip_text_from_json_response(raw_response)
332
+ error_message, response = self._parse_response(raw_response, parser)
333
+
334
+ if response is None:
335
+ messages.extend(
336
+ [
337
+ {"role": "assistant", "content": raw_response},
338
+ {
339
+ "role": "user",
340
+ "content": f"Try again. Your previous response was invalid and led to this error message: {error_message}",
341
+ },
342
+ ]
343
+ )
344
+ continue
345
+
346
+ if extra_validation:
347
+ extra_error_message = extra_validation(response)
348
+ if extra_error_message:
349
+ if self.support_structured_outputs:
350
+ # the raw response was a pydantic object, so we need to dump it to a string
351
+ raw_response = raw_response.model_dump_json()
352
+ elif not isinstance(raw_response, str):
353
+ raise ValueError(
354
+ "The response should be a string if the model does not support structured outputs."
355
+ )
356
+ messages.extend(
357
+ [
358
+ {"role": "assistant", "content": raw_response},
359
+ {
360
+ "role": "user",
361
+ "content": f"Try again. Your previous response was invalid and led to this error message: {extra_error_message}",
362
+ },
363
+ ]
364
+ )
365
+ continue
366
+ break
367
+
368
+ if debug_saver is not None:
369
+ debug_saver(formatted_prompt, kwargs, response)
370
+
371
+ return response
372
+
373
+ def _generate_hash(self, prompt: str) -> str:
374
+ hash_object = hashlib.sha256(prompt.encode())
375
+ return hash_object.hexdigest()
376
+
377
+ def _parse_response(
378
+ self, response: str, parser: MinimalPydanticOutputParser
379
+ ) -> Tuple[str, Dict[str, Any]]:
380
+ self.logger.info("Parsing JSON response: %s", response)
381
+ error_message = None
382
+ try:
383
+ response = parser.parse(response)
384
+ except Exception as e:
385
+ self.logger.error("Error parsing response: %s", e)
386
+ error_message = str(e)
387
+ response = None
388
+ return error_message, response
389
+
390
+ @staticmethod
391
+ def get_format_instructions(pydantic_object: Type[BaseModel]) -> str:
392
+ """
393
+ Generate format instructions for a Pydantic model's JSON output.
394
+
395
+ This function creates a string of instructions on how to format JSON output
396
+ based on the schema of a given Pydantic model. It's compatible with both
397
+ Pydantic v1 and v2.
398
+
399
+ Args:
400
+ pydantic_object (Type[BaseModel]): The Pydantic model class to generate instructions for.
401
+
402
+ Returns:
403
+ str: A string containing the format instructions.
404
+
405
+ Note:
406
+ This function is adapted from the LangChain framework.
407
+ Original source: https://github.com/langchain-ai/langchain
408
+ License: MIT (https://github.com/langchain-ai/langchain/blob/master/LICENSE)
409
+ """
410
+ _PYDANTIC_FORMAT_INSTRUCTIONS = """The output should be formatted as a JSON instance that conforms to the JSON schema below.
411
+
412
+ As an example, for the schema {{"properties": {{"foo": {{"title": "Foo", "description": "a list of strings", "type": "array", "items": {{"type": "string"}}}}}}, "required": ["foo"]}}
413
+ the object {{"foo": ["bar", "baz"]}} is a well-formatted instance of the schema. The object {{"properties": {{"foo": ["bar", "baz"]}}}} is not well-formatted.
414
+
415
+ Here is the output schema:
416
+ ```
417
+ {schema}
418
+ ```
419
+ """
420
+ schema = pydantic_object.model_json_schema().copy()
421
+
422
+ schema.pop("title", None)
423
+ schema.pop("type", None)
424
+
425
+ schema_str = json.dumps(schema, ensure_ascii=False)
426
+
427
+ return _PYDANTIC_FORMAT_INSTRUCTIONS.format(schema=schema_str)
@@ -0,0 +1,211 @@
1
+ import inspect
2
+ import re
3
+ from dataclasses import dataclass, field
4
+ from textwrap import dedent
5
+ from typing import Any, Callable, Dict, Optional, Type, get_type_hints
6
+
7
+
8
+ @dataclass
9
+ class Tool:
10
+ name: str
11
+ description: str
12
+ parameters: Dict[str, Any]
13
+ func: Callable[..., Any] = field(repr=False)
14
+
15
+ def to_dict(self) -> Dict[str, Any]:
16
+ """Convert the tool to a dictionary format compatible with Ollama."""
17
+ return {
18
+ "type": "function",
19
+ "function": {
20
+ "name": self.name,
21
+ "description": self.description,
22
+ "parameters": self.parameters,
23
+ "strict": True,
24
+ },
25
+ }
26
+
27
+ def execute(self, **kwargs) -> Any:
28
+ return self.func(**kwargs)
29
+
30
+
31
+ def _type_to_json_schema(type_hint: Type) -> Dict[str, Any]:
32
+ """Convert Python type hints to JSON Schema types."""
33
+ if type_hint == str:
34
+ return {"type": "string"}
35
+ elif type_hint == int:
36
+ return {"type": "integer"}
37
+ elif type_hint == float:
38
+ return {"type": "number"}
39
+ elif type_hint == bool:
40
+ return {"type": "boolean"}
41
+ elif type_hint == list or getattr(type_hint, "__origin__", None) == list:
42
+ item_type = Any
43
+ if hasattr(type_hint, "__args__"):
44
+ item_type = type_hint.__args__[0]
45
+ return {"type": "array", "items": _type_to_json_schema(item_type)}
46
+ elif type_hint == dict or getattr(type_hint, "__origin__", None) == dict:
47
+ return {"type": "object"}
48
+ elif hasattr(type_hint, "__origin__") and type_hint.__origin__ == Optional:
49
+ return _type_to_json_schema(type_hint.__args__[0])
50
+ else:
51
+ return {"type": "string"}
52
+
53
+
54
+ def _parse_docstring(docstring: str) -> tuple[str, dict[str, str]]:
55
+ """
56
+ Parse a docstring to extract the main description and parameter descriptions.
57
+
58
+ Args:
59
+ docstring: The function's docstring
60
+
61
+ Returns:
62
+ tuple: (main_description, parameter_descriptions)
63
+ """
64
+ if not docstring:
65
+ return "", {}
66
+
67
+ # Split docstring into sections
68
+ parts = re.split(r"\n\s*\n", dedent(docstring).strip())
69
+
70
+ # Get main description (first paragraph)
71
+ main_desc = parts[0].strip()
72
+
73
+ # Parse parameter descriptions
74
+ param_desc = {}
75
+ current_param = None
76
+ in_args_section = False
77
+
78
+ # Join all parts after the main description
79
+ remaining_text = "\n".join(parts[1:]) if len(parts) > 1 else ""
80
+
81
+ args_lines = []
82
+ for line in remaining_text.split("\n"):
83
+ line = line.rstrip()
84
+
85
+ # Check if we're entering the Args section
86
+ if line.lower().endswith("args:"):
87
+ in_args_section = True
88
+ continue
89
+
90
+ if not in_args_section:
91
+ continue
92
+
93
+ if line and not line.startswith(" "):
94
+ in_args_section = False
95
+ continue
96
+
97
+ args_lines.append(line)
98
+
99
+ # find parameter descriptions
100
+ args_lines = dedent("\n".join(args_lines)).split("\n")
101
+ for line in args_lines:
102
+ # Check for new parameter
103
+ if line and not line.startswith(" "):
104
+ # Look for parameter definition (param: description)
105
+ param_match = re.match(r"(\w+):\s*(.*)", line)
106
+ print(param_match)
107
+ if param_match:
108
+ current_param = param_match.group(1)
109
+ param_desc[current_param] = param_match.group(2)
110
+ # Add to existing parameter description
111
+ elif current_param and line:
112
+ param_desc[current_param] = param_desc[current_param] + " " + line.strip()
113
+
114
+ return main_desc, param_desc
115
+
116
+
117
+ def create_tool(
118
+ func: Callable[..., Any],
119
+ name: Optional[str] = None,
120
+ description: Optional[str] = None,
121
+ ) -> Tool:
122
+ """
123
+ Create a Tool instance from a Python function using its type hints and docstring.
124
+
125
+ Args:
126
+ func (Callable): The function to convert into a tool
127
+ name (Optional[str]): Optional custom name for the tool. If not provided, uses the function name
128
+ description (Optional[str]): Optional custom description. If not provided, uses the function's docstring
129
+
130
+ Returns:
131
+ Tool: A Tool instance representing the function
132
+
133
+ Example:
134
+ @create_tool
135
+ def get_weather(location: str, units: str = "celsius") -> str:
136
+ '''Get the weather for a specific location.
137
+
138
+ Args:
139
+ location: The city or location to get weather for
140
+ units: Temperature units (celsius or fahrenheit)
141
+ '''
142
+ # Function implementation
143
+ pass
144
+ """
145
+ # Get function metadata
146
+ func_name = name or func.__name__
147
+ func_doc = inspect.getdoc(func) or ""
148
+ func_desc = description or func_doc.split("\n\n")[0] if func_doc else func_name
149
+
150
+ # Parse docstring
151
+ main_desc, param_docs = _parse_docstring(func_doc)
152
+ func_desc = description or main_desc or func_name
153
+
154
+ # Get type hints
155
+ type_hints = get_type_hints(func)
156
+
157
+ # Get default values
158
+ signature = inspect.signature(func)
159
+
160
+ # Build parameters schema
161
+ parameters = {
162
+ "type": "object",
163
+ "properties": {},
164
+ "required": [],
165
+ "additionalProperties": False,
166
+ }
167
+
168
+ for param_name, param in signature.parameters.items():
169
+ if param_name == "self": # Skip self parameter for methods
170
+ continue
171
+
172
+ param_schema = _type_to_json_schema(type_hints.get(param_name, Any))
173
+
174
+ # Add description if available
175
+ if param_name in param_docs:
176
+ param_schema["description"] = param_docs[param_name]
177
+
178
+ # Handle default values
179
+ if param.default is not inspect.Parameter.empty:
180
+ param_schema["default"] = param.default
181
+ else:
182
+ parameters["required"].append(param_name)
183
+
184
+ parameters["properties"][param_name] = param_schema
185
+
186
+ return Tool(name=func_name, description=func_desc, parameters=parameters, func=func)
187
+
188
+
189
+ def tool(name: Optional[str] = None, description: Optional[str] = None) -> Callable:
190
+ """
191
+ Decorator to create a Tool from a function.
192
+
193
+ Args:
194
+ name (Optional[str]): Optional custom name for the tool
195
+ description (Optional[str]): Optional custom description
196
+
197
+ Returns:
198
+ Callable: Decorator function that creates a Tool
199
+
200
+ Example:
201
+ @tool(name="weather", description="Get weather information")
202
+ def get_weather(location: str, units: str = "celsius") -> str:
203
+ '''Get the weather for a specific location.'''
204
+ # Function implementation
205
+ pass
206
+ """
207
+
208
+ def decorator(func: Callable) -> Tool:
209
+ return create_tool(func, name=name, description=description)
210
+
211
+ return decorator