utcp 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.
- utcp/__init__.py +42 -0
- utcp/client/__init__.py +0 -0
- utcp/client/client_transport_interface.py +17 -0
- utcp/client/openapi_converter.py +301 -0
- utcp/client/tool_repositories/in_mem_tool_repository.py +47 -0
- utcp/client/tool_repository.py +101 -0
- utcp/client/tool_search_strategies/tag_search.py +67 -0
- utcp/client/tool_search_strategy.py +18 -0
- utcp/client/transport_interfaces/cli_transport.py +363 -0
- utcp/client/transport_interfaces/http_transport.py +310 -0
- utcp/client/transport_interfaces/mcp_transport.py +264 -0
- utcp/client/transport_interfaces/sse_transport.py +312 -0
- utcp/client/transport_interfaces/streamable_http_transport.py +332 -0
- utcp/client/transport_interfaces/text_transport.py +166 -0
- utcp/client/utcp_client.py +321 -0
- utcp/client/utcp_client_config.py +36 -0
- utcp/shared/__init__.py +0 -0
- utcp/shared/auth.py +39 -0
- utcp/shared/provider.py +194 -0
- utcp/shared/tool.py +108 -0
- utcp/shared/utcp_manual.py +21 -0
- utcp/version.py +16 -0
- utcp-0.1.0.dist-info/METADATA +585 -0
- utcp-0.1.0.dist-info/RECORD +27 -0
- utcp-0.1.0.dist-info/WHEEL +5 -0
- utcp-0.1.0.dist-info/licenses/LICENSE +373 -0
- utcp-0.1.0.dist-info/top_level.txt +1 -0
utcp/__init__.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Universal Tool Calling Protocol Core
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from utcp.shared.tool import (
|
|
6
|
+
Tool,
|
|
7
|
+
ToolInputOutputSchema,
|
|
8
|
+
)
|
|
9
|
+
|
|
10
|
+
from utcp.shared.provider import (
|
|
11
|
+
Provider,
|
|
12
|
+
HttpProvider,
|
|
13
|
+
CliProvider,
|
|
14
|
+
WebSocketProvider,
|
|
15
|
+
GRPCProvider,
|
|
16
|
+
GraphQLProvider,
|
|
17
|
+
TCPProvider,
|
|
18
|
+
UDPProvider,
|
|
19
|
+
StreamableHttpProvider,
|
|
20
|
+
SSEProvider,
|
|
21
|
+
WebRTCProvider,
|
|
22
|
+
MCPProvider,
|
|
23
|
+
TextProvider,
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
__all__ = [
|
|
27
|
+
"Tool",
|
|
28
|
+
"ToolInputOutputSchema",
|
|
29
|
+
"Provider",
|
|
30
|
+
"HttpProvider",
|
|
31
|
+
"CliProvider",
|
|
32
|
+
"WebSocketProvider",
|
|
33
|
+
"GRPCProvider",
|
|
34
|
+
"GraphQLProvider",
|
|
35
|
+
"TCPProvider",
|
|
36
|
+
"UDPProvider",
|
|
37
|
+
"StreamableHttpProvider",
|
|
38
|
+
"SSEProvider",
|
|
39
|
+
"WebRTCProvider",
|
|
40
|
+
"MCPProvider",
|
|
41
|
+
"TextProvider",
|
|
42
|
+
]
|
utcp/client/__init__.py
ADDED
|
File without changes
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from typing import Dict, Any, List
|
|
3
|
+
from utcp.shared.provider import Provider
|
|
4
|
+
from utcp.shared.tool import Tool
|
|
5
|
+
|
|
6
|
+
class ClientTransportInterface(ABC):
|
|
7
|
+
@abstractmethod
|
|
8
|
+
async def register_tool_provider(self, manual_provider: Provider) -> List[Tool]:
|
|
9
|
+
pass
|
|
10
|
+
|
|
11
|
+
@abstractmethod
|
|
12
|
+
async def deregister_tool_provider(self, manual_provider: Provider) -> None:
|
|
13
|
+
pass
|
|
14
|
+
|
|
15
|
+
@abstractmethod
|
|
16
|
+
async def call_tool(self, tool_name: str, arguments: Dict[str, Any], tool_provider: Provider) -> Any:
|
|
17
|
+
pass
|
|
@@ -0,0 +1,301 @@
|
|
|
1
|
+
import json
|
|
2
|
+
from typing import Any, Dict, List, Optional, Tuple
|
|
3
|
+
import sys
|
|
4
|
+
from utcp.shared.tool import Tool, ToolInputOutputSchema
|
|
5
|
+
from utcp.shared.utcp_manual import UtcpManual
|
|
6
|
+
from urllib.parse import urlparse
|
|
7
|
+
|
|
8
|
+
from utcp.shared.provider import HttpProvider
|
|
9
|
+
from utcp.shared.auth import Auth, ApiKeyAuth, BasicAuth, OAuth2Auth
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class OpenApiConverter:
|
|
13
|
+
"""Converts an OpenAPI JSON specification into a UtcpManual."""
|
|
14
|
+
|
|
15
|
+
def __init__(self, openapi_spec: Dict[str, Any], spec_url: Optional[str] = None, provider_name: Optional[str] = None):
|
|
16
|
+
self.spec = openapi_spec
|
|
17
|
+
self.spec_url = spec_url
|
|
18
|
+
# If provider_name is None then get the first word in spec.info.title
|
|
19
|
+
if provider_name is None:
|
|
20
|
+
title = openapi_spec.get("info", {}).get("title", "openapi_provider")
|
|
21
|
+
# Replace characters that are invalid for identifiers
|
|
22
|
+
invalid_chars = " -.,!?'\"\\/()[]{}#@$%^&*+=~`|;:<>"
|
|
23
|
+
self.provider_name = ''.join('_' if c in invalid_chars else c for c in title)
|
|
24
|
+
else:
|
|
25
|
+
self.provider_name = provider_name
|
|
26
|
+
|
|
27
|
+
def convert(self) -> UtcpManual:
|
|
28
|
+
"""Parses the OpenAPI specification and returns a UtcpManual."""
|
|
29
|
+
tools = []
|
|
30
|
+
servers = self.spec.get("servers")
|
|
31
|
+
if servers:
|
|
32
|
+
base_url = servers[0].get("url", "/")
|
|
33
|
+
elif self.spec_url:
|
|
34
|
+
parsed_url = urlparse(self.spec_url)
|
|
35
|
+
base_url = f"{parsed_url.scheme}://{parsed_url.netloc}"
|
|
36
|
+
else:
|
|
37
|
+
# Fallback if no server info and no spec URL is provided
|
|
38
|
+
base_url = "/"
|
|
39
|
+
print("No server info or spec URL provided. Using fallback base URL: ", base_url, file=sys.stderr)
|
|
40
|
+
|
|
41
|
+
for path, path_item in self.spec.get("paths", {}).items():
|
|
42
|
+
for method, operation in path_item.items():
|
|
43
|
+
if method.lower() in ['get', 'post', 'put', 'delete', 'patch']:
|
|
44
|
+
tool = self._create_tool(path, method, operation, base_url)
|
|
45
|
+
if tool:
|
|
46
|
+
tools.append(tool)
|
|
47
|
+
|
|
48
|
+
return UtcpManual(tools=tools)
|
|
49
|
+
|
|
50
|
+
def _resolve_ref(self, ref: str) -> Dict[str, Any]:
|
|
51
|
+
"""Resolves a local JSON reference."""
|
|
52
|
+
if not ref.startswith('#/'):
|
|
53
|
+
raise ValueError(f"External or non-local references are not supported: {ref}")
|
|
54
|
+
|
|
55
|
+
parts = ref[2:].split('/')
|
|
56
|
+
node = self.spec
|
|
57
|
+
for part in parts:
|
|
58
|
+
try:
|
|
59
|
+
node = node[part]
|
|
60
|
+
except (KeyError, TypeError):
|
|
61
|
+
raise ValueError(f"Reference not found: {ref}")
|
|
62
|
+
return node
|
|
63
|
+
|
|
64
|
+
def _resolve_schema(self, schema: Dict[str, Any]) -> Dict[str, Any]:
|
|
65
|
+
"""Recursively resolves all $refs in a schema object."""
|
|
66
|
+
if isinstance(schema, dict):
|
|
67
|
+
if "$ref" in schema:
|
|
68
|
+
resolved_ref = self._resolve_ref(schema["$ref"])
|
|
69
|
+
# The resolved reference could itself contain refs, so we recurse
|
|
70
|
+
return self._resolve_schema(resolved_ref)
|
|
71
|
+
|
|
72
|
+
# Resolve refs in nested properties
|
|
73
|
+
new_schema = {}
|
|
74
|
+
for key, value in schema.items():
|
|
75
|
+
new_schema[key] = self._resolve_schema(value)
|
|
76
|
+
return new_schema
|
|
77
|
+
|
|
78
|
+
if isinstance(schema, list):
|
|
79
|
+
return [self._resolve_schema(item) for item in schema]
|
|
80
|
+
|
|
81
|
+
return schema
|
|
82
|
+
|
|
83
|
+
def _extract_auth(self, operation: Dict[str, Any]) -> Optional[Auth]:
|
|
84
|
+
"""Extracts authentication information from OpenAPI operation and global security schemes."""
|
|
85
|
+
# First check for operation-level security requirements
|
|
86
|
+
security_requirements = operation.get("security", [])
|
|
87
|
+
|
|
88
|
+
# If no operation-level security, check global security requirements
|
|
89
|
+
if not security_requirements:
|
|
90
|
+
security_requirements = self.spec.get("security", [])
|
|
91
|
+
|
|
92
|
+
# If no security requirements, return None
|
|
93
|
+
if not security_requirements:
|
|
94
|
+
return None
|
|
95
|
+
|
|
96
|
+
# Get security schemes - support both OpenAPI 2.0 and 3.0
|
|
97
|
+
security_schemes = self._get_security_schemes()
|
|
98
|
+
|
|
99
|
+
# Process the first security requirement (most common case)
|
|
100
|
+
# Each security requirement is a dict with scheme name as key
|
|
101
|
+
for security_req in security_requirements:
|
|
102
|
+
for scheme_name, scopes in security_req.items():
|
|
103
|
+
if scheme_name in security_schemes:
|
|
104
|
+
scheme = security_schemes[scheme_name]
|
|
105
|
+
return self._create_auth_from_scheme(scheme, scheme_name)
|
|
106
|
+
|
|
107
|
+
return None
|
|
108
|
+
|
|
109
|
+
def _get_security_schemes(self) -> Dict[str, Any]:
|
|
110
|
+
"""Gets security schemes supporting both OpenAPI 2.0 and 3.0."""
|
|
111
|
+
# OpenAPI 3.0 format
|
|
112
|
+
if "components" in self.spec:
|
|
113
|
+
return self.spec.get("components", {}).get("securitySchemes", {})
|
|
114
|
+
|
|
115
|
+
# OpenAPI 2.0 format
|
|
116
|
+
return self.spec.get("securityDefinitions", {})
|
|
117
|
+
|
|
118
|
+
def _create_auth_from_scheme(self, scheme: Dict[str, Any], scheme_name: str) -> Optional[Auth]:
|
|
119
|
+
"""Creates an Auth object from an OpenAPI security scheme."""
|
|
120
|
+
scheme_type = scheme.get("type", "").lower()
|
|
121
|
+
|
|
122
|
+
if scheme_type == "apikey":
|
|
123
|
+
# For API key auth, use the parameter name from the OpenAPI spec
|
|
124
|
+
location = scheme.get("in", "header") # Default to header if not specified
|
|
125
|
+
param_name = scheme.get("name", "Authorization") # Default name
|
|
126
|
+
return ApiKeyAuth(
|
|
127
|
+
api_key=f"${{{self.provider_name.upper()}_API_KEY}}", # Placeholder for environment variable
|
|
128
|
+
var_name=param_name,
|
|
129
|
+
location=location
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
elif scheme_type == "basic":
|
|
133
|
+
# OpenAPI 2.0 format: type: basic
|
|
134
|
+
return BasicAuth(
|
|
135
|
+
username=f"${{{self.provider_name.upper()}_USERNAME}}",
|
|
136
|
+
password=f"${{{self.provider_name.upper()}_PASSWORD}}"
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
elif scheme_type == "http":
|
|
140
|
+
# OpenAPI 3.0 format: type: http with scheme
|
|
141
|
+
http_scheme = scheme.get("scheme", "").lower()
|
|
142
|
+
if http_scheme == "basic":
|
|
143
|
+
# For basic auth, use conventional environment variable names
|
|
144
|
+
return BasicAuth(
|
|
145
|
+
username=f"${{{self.provider_name.upper()}_USERNAME}}",
|
|
146
|
+
password=f"${{{self.provider_name.upper()}_PASSWORD}}"
|
|
147
|
+
)
|
|
148
|
+
elif http_scheme == "bearer":
|
|
149
|
+
# Treat bearer tokens as API keys
|
|
150
|
+
return ApiKeyAuth(
|
|
151
|
+
api_key=f"Bearer ${{{self.provider_name.upper()}_API_KEY}}",
|
|
152
|
+
var_name="Authorization",
|
|
153
|
+
location="header"
|
|
154
|
+
)
|
|
155
|
+
|
|
156
|
+
elif scheme_type == "oauth2":
|
|
157
|
+
# Handle both OpenAPI 2.0 and 3.0 OAuth2 formats
|
|
158
|
+
flows = scheme.get("flows", {})
|
|
159
|
+
|
|
160
|
+
# OpenAPI 3.0 format
|
|
161
|
+
if flows:
|
|
162
|
+
for flow_type, flow_config in flows.items():
|
|
163
|
+
# Support both old and new flow names
|
|
164
|
+
if flow_type in ["authorizationCode", "accessCode", "clientCredentials", "application"]:
|
|
165
|
+
token_url = flow_config.get("tokenUrl")
|
|
166
|
+
if token_url:
|
|
167
|
+
return OAuth2Auth(
|
|
168
|
+
token_url=token_url,
|
|
169
|
+
client_id=f"${{{self.provider_name.upper()}_CLIENT_ID}}",
|
|
170
|
+
client_secret=f"${{{self.provider_name.upper()}_CLIENT_SECRET}}",
|
|
171
|
+
scope=" ".join(flow_config.get("scopes", {}).keys()) or None
|
|
172
|
+
)
|
|
173
|
+
|
|
174
|
+
# OpenAPI 2.0 format (flows directly in scheme)
|
|
175
|
+
else:
|
|
176
|
+
flow_type = scheme.get("flow", "")
|
|
177
|
+
token_url = scheme.get("tokenUrl")
|
|
178
|
+
if token_url and flow_type in ["accessCode", "application", "clientCredentials"]:
|
|
179
|
+
return OAuth2Auth(
|
|
180
|
+
token_url=token_url,
|
|
181
|
+
client_id=f"${{{self.provider_name.upper()}_CLIENT_ID}}",
|
|
182
|
+
client_secret=f"${{{self.provider_name.upper()}_CLIENT_SECRET}}",
|
|
183
|
+
scope=" ".join(scheme.get("scopes", {}).keys()) or None
|
|
184
|
+
)
|
|
185
|
+
|
|
186
|
+
return None
|
|
187
|
+
|
|
188
|
+
def _create_tool(self, path: str, method: str, operation: Dict[str, Any], base_url: str) -> Optional[Tool]:
|
|
189
|
+
"""Creates a Tool object from an OpenAPI operation."""
|
|
190
|
+
operation_id = operation.get("operationId")
|
|
191
|
+
if not operation_id:
|
|
192
|
+
return None
|
|
193
|
+
|
|
194
|
+
description = operation.get("summary") or operation.get("description", "")
|
|
195
|
+
tags = operation.get("tags", [])
|
|
196
|
+
|
|
197
|
+
inputs, header_fields, body_field = self._extract_inputs(operation)
|
|
198
|
+
outputs = self._extract_outputs(operation)
|
|
199
|
+
auth = self._extract_auth(operation)
|
|
200
|
+
|
|
201
|
+
provider_name = self.spec.get("info", {}).get("title", "openapi_provider")
|
|
202
|
+
|
|
203
|
+
# Combine base URL and path, ensuring no double slashes
|
|
204
|
+
full_url = base_url.rstrip('/') + '/' + path.lstrip('/')
|
|
205
|
+
|
|
206
|
+
provider = HttpProvider(
|
|
207
|
+
name=provider_name,
|
|
208
|
+
provider_type="http",
|
|
209
|
+
http_method=method.upper(),
|
|
210
|
+
url=full_url,
|
|
211
|
+
body_field=body_field if body_field else None,
|
|
212
|
+
header_fields=header_fields if header_fields else None,
|
|
213
|
+
auth=auth
|
|
214
|
+
)
|
|
215
|
+
|
|
216
|
+
return Tool(
|
|
217
|
+
name=operation_id,
|
|
218
|
+
description=description,
|
|
219
|
+
inputs=inputs,
|
|
220
|
+
outputs=outputs,
|
|
221
|
+
tags=tags,
|
|
222
|
+
tool_provider=provider
|
|
223
|
+
)
|
|
224
|
+
|
|
225
|
+
def _extract_inputs(self, operation: Dict[str, Any]) -> Tuple[ToolInputOutputSchema, List[str], Optional[str]]:
|
|
226
|
+
"""Extracts input schema, header fields, and body field from an OpenAPI operation."""
|
|
227
|
+
properties = {}
|
|
228
|
+
required = []
|
|
229
|
+
header_fields = []
|
|
230
|
+
body_field = None
|
|
231
|
+
|
|
232
|
+
# Handle parameters (query, header, path, cookie)
|
|
233
|
+
for param in operation.get("parameters", []):
|
|
234
|
+
param = self._resolve_schema(param)
|
|
235
|
+
param_name = param.get("name")
|
|
236
|
+
if not param_name:
|
|
237
|
+
continue
|
|
238
|
+
|
|
239
|
+
if param.get("in") == "header":
|
|
240
|
+
header_fields.append(param_name)
|
|
241
|
+
|
|
242
|
+
schema = self._resolve_schema(param.get("schema", {}))
|
|
243
|
+
properties[param_name] = {
|
|
244
|
+
"type": schema.get("type", "string"),
|
|
245
|
+
"description": param.get("description", ""),
|
|
246
|
+
**schema
|
|
247
|
+
}
|
|
248
|
+
if param.get("required"):
|
|
249
|
+
required.append(param_name)
|
|
250
|
+
|
|
251
|
+
# Handle request body
|
|
252
|
+
request_body = operation.get("requestBody")
|
|
253
|
+
if request_body:
|
|
254
|
+
resolved_body = self._resolve_schema(request_body)
|
|
255
|
+
content = resolved_body.get("content", {})
|
|
256
|
+
json_schema = content.get("application/json", {}).get("schema")
|
|
257
|
+
if json_schema:
|
|
258
|
+
# Add a single 'body' field to represent the request body
|
|
259
|
+
body_field = "body"
|
|
260
|
+
properties[body_field] = {
|
|
261
|
+
"description": resolved_body.get("description", "Request body"),
|
|
262
|
+
**self._resolve_schema(json_schema)
|
|
263
|
+
}
|
|
264
|
+
if resolved_body.get("required"):
|
|
265
|
+
required.append(body_field)
|
|
266
|
+
|
|
267
|
+
schema = ToolInputOutputSchema(properties=properties, required=required if required else None)
|
|
268
|
+
return schema, header_fields, body_field
|
|
269
|
+
|
|
270
|
+
def _extract_outputs(self, operation: Dict[str, Any]) -> ToolInputOutputSchema:
|
|
271
|
+
"""Extracts the output schema from an OpenAPI operation, resolving refs."""
|
|
272
|
+
success_response = operation.get("responses", {}).get("200") or operation.get("responses", {}).get("201")
|
|
273
|
+
if not success_response:
|
|
274
|
+
return ToolInputOutputSchema()
|
|
275
|
+
|
|
276
|
+
resolved_response = self._resolve_schema(success_response)
|
|
277
|
+
content = resolved_response.get("content", {})
|
|
278
|
+
json_schema = content.get("application/json", {}).get("schema")
|
|
279
|
+
|
|
280
|
+
if not json_schema:
|
|
281
|
+
return ToolInputOutputSchema()
|
|
282
|
+
|
|
283
|
+
resolved_json_schema = self._resolve_schema(json_schema)
|
|
284
|
+
schema_args = {
|
|
285
|
+
"type": resolved_json_schema.get("type", "object"),
|
|
286
|
+
"properties": resolved_json_schema.get("properties", {}),
|
|
287
|
+
"required": resolved_json_schema.get("required"),
|
|
288
|
+
"description": resolved_json_schema.get("description"),
|
|
289
|
+
"title": resolved_json_schema.get("title"),
|
|
290
|
+
}
|
|
291
|
+
|
|
292
|
+
# Handle array item types
|
|
293
|
+
if schema_args["type"] == "array" and "items" in resolved_json_schema:
|
|
294
|
+
schema_args["items"] = resolved_json_schema.get("items")
|
|
295
|
+
|
|
296
|
+
# Handle additional schema attributes
|
|
297
|
+
for attr in ["enum", "minimum", "maximum", "format"]:
|
|
298
|
+
if attr in resolved_json_schema:
|
|
299
|
+
schema_args[attr] = resolved_json_schema.get(attr)
|
|
300
|
+
|
|
301
|
+
return ToolInputOutputSchema(**schema_args)
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
from typing import List, Dict, Tuple, Optional
|
|
2
|
+
from utcp.shared.provider import Provider
|
|
3
|
+
from utcp.shared.tool import Tool
|
|
4
|
+
from utcp.client.tool_repository import ToolRepository
|
|
5
|
+
|
|
6
|
+
class InMemToolRepository(ToolRepository):
|
|
7
|
+
tools: List[Tool] = []
|
|
8
|
+
tool_per_provider: Dict[str, Tuple[Provider, List[Tool]]] = {}
|
|
9
|
+
|
|
10
|
+
async def save_provider_with_tools(self, provider: Provider, tools: List[Tool]) -> None:
|
|
11
|
+
self.tools.extend(tools)
|
|
12
|
+
self.tool_per_provider[provider.name] = (provider, tools)
|
|
13
|
+
|
|
14
|
+
async def remove_provider(self, provider_name: str) -> None:
|
|
15
|
+
if provider_name not in self.tool_per_provider:
|
|
16
|
+
raise ValueError(f"Provider '{provider_name}' not found")
|
|
17
|
+
tools_to_remove = self.tool_per_provider[provider_name][1]
|
|
18
|
+
self.tools = [tool for tool in self.tools if tool not in tools_to_remove]
|
|
19
|
+
self.tool_per_provider.pop(provider_name, None)
|
|
20
|
+
|
|
21
|
+
async def remove_tool(self, tool_name: str) -> None:
|
|
22
|
+
provider_name = tool_name.split(".")[0]
|
|
23
|
+
if provider_name not in self.tool_per_provider:
|
|
24
|
+
raise ValueError(f"Provider '{provider_name}' not found")
|
|
25
|
+
new_tools = [tool for tool in self.tools if tool.name != tool_name]
|
|
26
|
+
if len(new_tools) == len(self.tools):
|
|
27
|
+
raise ValueError(f"Tool '{tool_name}' not found")
|
|
28
|
+
self.tools = new_tools
|
|
29
|
+
self.tool_per_provider[provider_name][1] = [tool for tool in self.tool_per_provider[provider_name][1] if tool.name != tool_name]
|
|
30
|
+
|
|
31
|
+
async def get_tool(self, tool_name: str) -> Optional[Tool]:
|
|
32
|
+
for tool in self.tools:
|
|
33
|
+
if tool.name == tool_name:
|
|
34
|
+
return tool
|
|
35
|
+
return None
|
|
36
|
+
|
|
37
|
+
async def get_tools(self) -> List[Tool]:
|
|
38
|
+
return self.tools
|
|
39
|
+
|
|
40
|
+
async def get_tools_by_provider(self, provider_name: str) -> Optional[List[Tool]]:
|
|
41
|
+
return self.tool_per_provider.get(provider_name, (None, None))[1]
|
|
42
|
+
|
|
43
|
+
async def get_provider(self, provider_name: str) -> Optional[Provider]:
|
|
44
|
+
return self.tool_per_provider.get(provider_name, (None, None))[0]
|
|
45
|
+
|
|
46
|
+
async def get_providers(self) -> List[Provider]:
|
|
47
|
+
return [provider for provider, _ in self.tool_per_provider.values()]
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from typing import List, Dict, Any, Optional
|
|
3
|
+
from utcp.shared.provider import Provider
|
|
4
|
+
from utcp.shared.tool import Tool
|
|
5
|
+
|
|
6
|
+
class ToolRepository(ABC):
|
|
7
|
+
@abstractmethod
|
|
8
|
+
async def save_provider_with_tools(self, provider: Provider, tools: List[Tool]) -> None:
|
|
9
|
+
"""
|
|
10
|
+
Save a provider and its tools in the repository.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
provider: The provider to save.
|
|
14
|
+
tools: The tools associated with the provider.
|
|
15
|
+
"""
|
|
16
|
+
pass
|
|
17
|
+
|
|
18
|
+
@abstractmethod
|
|
19
|
+
async def remove_provider(self, provider_name: str) -> None:
|
|
20
|
+
"""
|
|
21
|
+
Remove a provider and its tools from the repository.
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
provider_name: The name of the provider to remove.
|
|
25
|
+
|
|
26
|
+
Raises:
|
|
27
|
+
ValueError: If the provider is not found.
|
|
28
|
+
"""
|
|
29
|
+
pass
|
|
30
|
+
|
|
31
|
+
@abstractmethod
|
|
32
|
+
async def remove_tool(self, tool_name: str) -> None:
|
|
33
|
+
"""
|
|
34
|
+
Remove a tool from the repository.
|
|
35
|
+
|
|
36
|
+
Args:
|
|
37
|
+
tool_name: The name of the tool to remove.
|
|
38
|
+
|
|
39
|
+
Raises:
|
|
40
|
+
ValueError: If the tool is not found.
|
|
41
|
+
"""
|
|
42
|
+
pass
|
|
43
|
+
|
|
44
|
+
@abstractmethod
|
|
45
|
+
async def get_tool(self, tool_name: str) -> Optional[Tool]:
|
|
46
|
+
"""
|
|
47
|
+
Get a tool from the repository.
|
|
48
|
+
|
|
49
|
+
Args:
|
|
50
|
+
tool_name: The name of the tool to retrieve.
|
|
51
|
+
|
|
52
|
+
Returns:
|
|
53
|
+
The tool if found, otherwise None.
|
|
54
|
+
"""
|
|
55
|
+
pass
|
|
56
|
+
|
|
57
|
+
@abstractmethod
|
|
58
|
+
async def get_tools(self) -> List[Tool]:
|
|
59
|
+
"""
|
|
60
|
+
Get all tools from the repository.
|
|
61
|
+
|
|
62
|
+
Returns:
|
|
63
|
+
A list of tools.
|
|
64
|
+
"""
|
|
65
|
+
pass
|
|
66
|
+
|
|
67
|
+
@abstractmethod
|
|
68
|
+
async def get_tools_by_provider(self, provider_name: str) -> Optional[List[Tool]]:
|
|
69
|
+
"""
|
|
70
|
+
Get tools associated with a specific provider.
|
|
71
|
+
|
|
72
|
+
Args:
|
|
73
|
+
provider_name: The name of the provider.
|
|
74
|
+
|
|
75
|
+
Returns:
|
|
76
|
+
A list of tools associated with the provider, or None if the provider is not found.
|
|
77
|
+
"""
|
|
78
|
+
pass
|
|
79
|
+
|
|
80
|
+
@abstractmethod
|
|
81
|
+
async def get_provider(self, provider_name: str) -> Optional[Provider]:
|
|
82
|
+
"""
|
|
83
|
+
Get a provider from the repository.
|
|
84
|
+
|
|
85
|
+
Args:
|
|
86
|
+
provider_name: The name of the provider to retrieve.
|
|
87
|
+
|
|
88
|
+
Returns:
|
|
89
|
+
The provider if found, otherwise None.
|
|
90
|
+
"""
|
|
91
|
+
pass
|
|
92
|
+
|
|
93
|
+
@abstractmethod
|
|
94
|
+
async def get_providers(self) -> List[Provider]:
|
|
95
|
+
"""
|
|
96
|
+
Get all providers from the repository.
|
|
97
|
+
|
|
98
|
+
Returns:
|
|
99
|
+
A list of providers.
|
|
100
|
+
"""
|
|
101
|
+
pass
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
from utcp.client.tool_search_strategy import ToolSearchStrategy
|
|
2
|
+
from typing import List, Dict, Tuple
|
|
3
|
+
from utcp.shared.tool import Tool
|
|
4
|
+
from utcp.client.tool_repository import ToolRepository
|
|
5
|
+
import re
|
|
6
|
+
import asyncio
|
|
7
|
+
|
|
8
|
+
class TagSearchStrategy(ToolSearchStrategy):
|
|
9
|
+
|
|
10
|
+
def __init__(self, tool_repository: ToolRepository, description_weight: float = 0.3):
|
|
11
|
+
self.tool_repository = tool_repository
|
|
12
|
+
# Weight for description words vs explicit tags (explicit tags have weight of 1.0)
|
|
13
|
+
self.description_weight = description_weight
|
|
14
|
+
|
|
15
|
+
async def search_tools(self, query: str, limit: int = 10) -> List[Tool]:
|
|
16
|
+
"""
|
|
17
|
+
Return tools ordered by tag occurrences in the query.
|
|
18
|
+
|
|
19
|
+
Uses both explicit tags and words from tool descriptions (with less weight).
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
query: The search query string
|
|
23
|
+
limit: Maximum number of tools to return
|
|
24
|
+
|
|
25
|
+
Returns:
|
|
26
|
+
List of tools ordered by relevance to the query
|
|
27
|
+
"""
|
|
28
|
+
# Normalize query to lowercase and split into words
|
|
29
|
+
query_lower = query.lower()
|
|
30
|
+
# Extract words from the query, filtering out non-word characters
|
|
31
|
+
query_words = set(re.findall(r'\w+', query_lower))
|
|
32
|
+
|
|
33
|
+
# Get all tools (using asyncio to run the coroutine)
|
|
34
|
+
tools = await self.tool_repository.get_tools()
|
|
35
|
+
|
|
36
|
+
# Calculate scores for each tool
|
|
37
|
+
tool_scores: List[Tuple[Tool, float]] = []
|
|
38
|
+
|
|
39
|
+
for tool in tools:
|
|
40
|
+
score = 0.0
|
|
41
|
+
|
|
42
|
+
# Score from explicit tags (weight 1.0)
|
|
43
|
+
for tag in tool.tags:
|
|
44
|
+
tag_lower = tag.lower()
|
|
45
|
+
# Check if the tag appears in the query
|
|
46
|
+
if tag_lower in query_lower:
|
|
47
|
+
score += 1.0
|
|
48
|
+
# Also check if the tag words match query words
|
|
49
|
+
tag_words = set(re.findall(r'\w+', tag_lower))
|
|
50
|
+
for word in tag_words:
|
|
51
|
+
if word in query_words:
|
|
52
|
+
score += self.description_weight # Partial match for tag words
|
|
53
|
+
|
|
54
|
+
# Score from description (with lower weight)
|
|
55
|
+
if tool.description:
|
|
56
|
+
description_words = set(re.findall(r'\w+', tool.description.lower()))
|
|
57
|
+
for word in description_words:
|
|
58
|
+
if word in query_words and len(word) > 2: # Only consider words with length > 2
|
|
59
|
+
score += self.description_weight
|
|
60
|
+
|
|
61
|
+
tool_scores.append((tool, score))
|
|
62
|
+
|
|
63
|
+
# Sort tools by score in descending order
|
|
64
|
+
sorted_tools = [tool for tool, score in sorted(tool_scores, key=lambda x: x[1], reverse=True)]
|
|
65
|
+
|
|
66
|
+
# Return up to 'limit' tools
|
|
67
|
+
return sorted_tools[:limit]
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from typing import List
|
|
3
|
+
from utcp.shared.tool import Tool
|
|
4
|
+
|
|
5
|
+
class ToolSearchStrategy(ABC):
|
|
6
|
+
@abstractmethod
|
|
7
|
+
async def search_tools(self, query: str, limit: int = 10) -> List[Tool]:
|
|
8
|
+
"""
|
|
9
|
+
Search for tools relevant to the query.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
query: The search query.
|
|
13
|
+
limit: The maximum number of tools to return. 0 for no limit.
|
|
14
|
+
|
|
15
|
+
Returns:
|
|
16
|
+
A list of tools that match the search query.
|
|
17
|
+
"""
|
|
18
|
+
pass
|