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:
2026-05-04 00:09:41 +03:00
parent d292d32afa
commit 6bbbefaa6f
4 changed files with 1170 additions and 274 deletions

View File

@@ -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)