feat: refactor APIClient and ToolRegistry to classes
- api_client.py: APIClient class with ProviderConfig types, chat_stream() returns str instead of printing, module-level wrappers for backward compat - tools.py: ToolRegistry class with instance methods, pure Python list_files (os.walk) and search_in_files (re.search) replacing find/grep for Windows compat, English tool descriptions and error messages, module-level wrappers for backward compat - 74 new tests (21 api_client + 53 tools) Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -1,26 +1,157 @@
|
||||
"""Universal API client supporting multiple OpenAI-compatible providers."""
|
||||
|
||||
import httpx
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Iterator
|
||||
|
||||
import httpx
|
||||
|
||||
from meeting_room.models import ProviderConfig
|
||||
|
||||
|
||||
PROVIDERS = {}
|
||||
class APIClient:
|
||||
"""Stateful client that holds provider configs and makes LLM calls.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
providers:
|
||||
Mapping of provider name to ``ProviderConfig`` (base_url + api_key).
|
||||
"""
|
||||
|
||||
def __init__(self, providers: dict[str, ProviderConfig]) -> None:
|
||||
self.providers: dict[str, ProviderConfig] = providers
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _get_client(self, provider_name: str) -> httpx.Client:
|
||||
"""Create a short-lived ``httpx.Client`` for *provider_name*."""
|
||||
p = self.providers.get(provider_name)
|
||||
if not p:
|
||||
raise ValueError(f"Unknown provider: {provider_name}")
|
||||
return httpx.Client(
|
||||
base_url=p.base_url,
|
||||
headers={"Authorization": f"Bearer {p.api_key}"},
|
||||
timeout=120.0,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public API
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def chat(
|
||||
self,
|
||||
provider: str,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
temperature: float = 0.7,
|
||||
tools: list[dict] | None = None,
|
||||
) -> dict:
|
||||
"""Send a non-streaming chat completion request.
|
||||
|
||||
Returns a dict with keys ``content`` (str) and ``tool_calls``
|
||||
(list | None).
|
||||
"""
|
||||
client = self._get_client(provider)
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
}
|
||||
if tools:
|
||||
payload["tools"] = tools
|
||||
payload["tool_choice"] = "auto"
|
||||
|
||||
resp = client.post("/chat/completions", json=payload)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
message = data["choices"][0]["message"]
|
||||
|
||||
result: dict = {
|
||||
"content": message.get("content") or "",
|
||||
"tool_calls": None,
|
||||
}
|
||||
|
||||
if message.get("tool_calls"):
|
||||
result["tool_calls"] = [
|
||||
{
|
||||
"id": tc["id"],
|
||||
"name": tc["function"]["name"],
|
||||
"arguments": tc["function"]["arguments"],
|
||||
}
|
||||
for tc in message["tool_calls"]
|
||||
]
|
||||
|
||||
return result
|
||||
|
||||
def chat_stream(
|
||||
self,
|
||||
provider: str,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
temperature: float = 0.7,
|
||||
) -> str:
|
||||
"""Send a streaming chat completion request.
|
||||
|
||||
Accumulates all content deltas and **returns** the full text.
|
||||
Does NOT print to stdout.
|
||||
"""
|
||||
client = self._get_client(provider)
|
||||
payload: dict = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"stream": True,
|
||||
}
|
||||
full_text = ""
|
||||
with client.stream("POST", "/chat/completions", json=payload) as resp:
|
||||
resp.raise_for_status()
|
||||
for line in resp.iter_lines():
|
||||
if not line or not line.startswith("data: "):
|
||||
continue
|
||||
data = line[6:]
|
||||
if data == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk = json.loads(data)
|
||||
delta = chunk["choices"][0]["delta"]
|
||||
if "content" in delta and delta["content"]:
|
||||
full_text += delta["content"]
|
||||
except (json.JSONDecodeError, KeyError, IndexError):
|
||||
continue
|
||||
return full_text
|
||||
|
||||
|
||||
def init_providers(config: dict):
|
||||
global PROVIDERS
|
||||
PROVIDERS = config.get("providers", {})
|
||||
# ======================================================================
|
||||
# Module-level backward-compatible wrappers
|
||||
# ======================================================================
|
||||
|
||||
_default_client: APIClient | None = None
|
||||
|
||||
|
||||
def _get_client(provider_name: str) -> httpx.Client:
|
||||
p = PROVIDERS.get(provider_name)
|
||||
if not p:
|
||||
raise ValueError(f"Unknown provider: {provider_name}")
|
||||
return httpx.Client(
|
||||
base_url=p["base_url"],
|
||||
headers={"Authorization": f"Bearer {p['api_key']}"},
|
||||
timeout=120.0,
|
||||
)
|
||||
def init_providers(config: dict) -> None:
|
||||
"""Initialise the module-level default ``APIClient``.
|
||||
|
||||
*config* is the raw dict produced by ``yaml.safe_load()`` — it must
|
||||
contain a ``providers`` key whose value maps provider names to
|
||||
``{base_url, api_key}`` dicts.
|
||||
"""
|
||||
global _default_client
|
||||
raw: dict = config.get("providers", {})
|
||||
providers: dict[str, ProviderConfig] = {
|
||||
name: ProviderConfig(**vals) for name, vals in raw.items()
|
||||
}
|
||||
_default_client = APIClient(providers)
|
||||
|
||||
|
||||
def _require_default() -> APIClient:
|
||||
if _default_client is None:
|
||||
raise RuntimeError(
|
||||
"init_providers() must be called before using module-level chat/chat_stream"
|
||||
)
|
||||
return _default_client
|
||||
|
||||
|
||||
def chat(
|
||||
@@ -30,36 +161,8 @@ def chat(
|
||||
temperature: float = 0.7,
|
||||
tools: list[dict] | None = None,
|
||||
) -> dict:
|
||||
client = _get_client(provider)
|
||||
payload = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
}
|
||||
if tools:
|
||||
payload["tools"] = tools
|
||||
payload["tool_choice"] = "auto"
|
||||
|
||||
resp = client.post("/chat/completions", json=payload)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
message = data["choices"][0]["message"]
|
||||
|
||||
result = {
|
||||
"content": message.get("content") or "",
|
||||
"tool_calls": None,
|
||||
}
|
||||
|
||||
if message.get("tool_calls"):
|
||||
result["tool_calls"] = []
|
||||
for tc in message["tool_calls"]:
|
||||
result["tool_calls"].append({
|
||||
"id": tc["id"],
|
||||
"name": tc["function"]["name"],
|
||||
"arguments": tc["function"]["arguments"],
|
||||
})
|
||||
|
||||
return result
|
||||
"""Module-level wrapper that delegates to the default ``APIClient``."""
|
||||
return _require_default().chat(provider, model, messages, temperature, tools)
|
||||
|
||||
|
||||
def chat_stream(
|
||||
@@ -68,29 +171,8 @@ def chat_stream(
|
||||
messages: list[dict],
|
||||
temperature: float = 0.7,
|
||||
) -> str:
|
||||
client = _get_client(provider)
|
||||
payload = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"stream": True,
|
||||
}
|
||||
full_text = ""
|
||||
with client.stream("POST", "/chat/completions", json=payload) as resp:
|
||||
resp.raise_for_status()
|
||||
for line in resp.iter_lines():
|
||||
if not line or not line.startswith("data: "):
|
||||
continue
|
||||
data = line[6:]
|
||||
if data == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk = json.loads(data)
|
||||
delta = chunk["choices"][0]["delta"]
|
||||
if "content" in delta and delta["content"]:
|
||||
print(delta["content"], end="", flush=True)
|
||||
full_text += delta["content"]
|
||||
except (json.JSONDecodeError, KeyError, IndexError):
|
||||
continue
|
||||
print()
|
||||
return full_text
|
||||
"""Module-level wrapper that delegates to the default ``APIClient``.
|
||||
|
||||
Returns the accumulated text — does NOT print to stdout.
|
||||
"""
|
||||
return _require_default().chat_stream(provider, model, messages, temperature)
|
||||
Reference in New Issue
Block a user