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.
- llm_interface/__init__.py +6 -0
- llm_interface/anthropic.py +194 -0
- llm_interface/llm_config.py +117 -0
- llm_interface/llm_interface.py +427 -0
- llm_interface/llm_tool.py +211 -0
- llm_interface/openai.py +162 -0
- llm_interface/pydantic_output_parser.py +99 -0
- llm_interface/remote_ollama.py +187 -0
- llm_interface/ssh.py +164 -0
- llm_interface/testing/__init__.py +0 -0
- llm_interface/testing/helpers.py +32 -0
- llm_interface/testing/mock_llm.py +84 -0
- llm_interface/utils.py +48 -0
- llm_interface-0.1.0.dist-info/LICENSE +201 -0
- llm_interface-0.1.0.dist-info/METADATA +177 -0
- llm_interface-0.1.0.dist-info/RECORD +17 -0
- llm_interface-0.1.0.dist-info/WHEEL +4 -0
|
@@ -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
|