diff --git a/pyproject.toml b/pyproject.toml
index 58af9ce..088021d 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "filthy-tool-fixer"
-version = "0.3.0"
+version = "0.3.1"
requires-python = ">=3.11"
dependencies = [
"fastapi>=0.109.0",
diff --git a/src/filthy_tool_fixer/backends/ollama.py b/src/filthy_tool_fixer/backends/ollama.py
index 727757d..4c7b1bd 100644
--- a/src/filthy_tool_fixer/backends/ollama.py
+++ b/src/filthy_tool_fixer/backends/ollama.py
@@ -90,11 +90,17 @@ async def chat_completion_stream(
json=payload,
timeout=httpx.Timeout(effective_timeout, connect=10.0),
) as resp:
- resp.raise_for_status()
+ try:
+ resp.raise_for_status()
+ except httpx.HTTPStatusError as e:
+ error = {"error": {"message": f"Backend returned {e.response.status_code}", "type": "backend_error"}}
+ yield f"data: {json.dumps(error)}\n\n".encode()
+ yield b"data: [DONE]\n\n"
+ return
try:
async for line in resp.aiter_lines():
if line:
- yield (line + "\n").encode()
+ yield (line + "\n\n").encode()
except asyncio.CancelledError:
log.info("stream_cancelled_by_client")
raise
diff --git a/src/filthy_tool_fixer/backends/vllm.py b/src/filthy_tool_fixer/backends/vllm.py
index e1121a2..e280923 100644
--- a/src/filthy_tool_fixer/backends/vllm.py
+++ b/src/filthy_tool_fixer/backends/vllm.py
@@ -2,6 +2,7 @@
from __future__ import annotations
+import asyncio
import json
from typing import Any, AsyncIterator
@@ -28,12 +29,12 @@ def build_tool_call_schema(tools: list[ToolDefinition]) -> dict[str, Any]:
tool_schemas = []
for tool in tools:
func = tool.function
- # Each tool call has name + arguments matching the tool's parameter schema
+ # Each tool call has name + arguments as a JSON string
tool_schema: dict[str, Any] = {
"type": "object",
"properties": {
"name": {"type": "string", "const": func.name},
- "arguments": func.parameters if func.parameters else {"type": "object"},
+ "arguments": {"type": "string"},
},
"required": ["name", "arguments"],
"additionalProperties": False,
@@ -87,7 +88,8 @@ async def chat_completion(
timeout: float | None = None,
keep_alive: str | None = None,
) -> ChatCompletionResponse:
- assert self._client is not None
+ if self._client is None:
+ raise RuntimeError("VLLMAdapter not started — call startup() first")
payload = self._build_payload(request, keep_alive)
payload["stream"] = False
@@ -106,7 +108,8 @@ async def chat_completion_stream(
timeout: float | None = None,
keep_alive: str | None = None,
) -> AsyncIterator[bytes]:
- assert self._client is not None
+ if self._client is None:
+ raise RuntimeError("VLLMAdapter not started — call startup() first")
payload = self._build_payload(request, keep_alive)
payload["stream"] = True
@@ -117,14 +120,25 @@ async def chat_completion_stream(
json=payload,
timeout=httpx.Timeout(effective_timeout, connect=10.0),
) as resp:
- resp.raise_for_status()
- async for line in resp.aiter_lines():
- if line:
- yield (line + "\n").encode()
+ try:
+ resp.raise_for_status()
+ except httpx.HTTPStatusError as e:
+ error = {"error": {"message": f"Backend returned {e.response.status_code}", "type": "backend_error"}}
+ yield f"data: {json.dumps(error)}\n\n".encode()
+ yield b"data: [DONE]\n\n"
+ return
+ try:
+ async for line in resp.aiter_lines():
+ if line:
+ yield (line + "\n\n").encode()
+ except asyncio.CancelledError:
+ log.info("stream_cancelled_by_client")
+ raise
async def health_check(self) -> bool:
try:
- assert self._client is not None
+ if self._client is None:
+ return False
resp = await self._client.get("/health", timeout=5.0)
return resp.status_code == 200
except Exception:
diff --git a/src/filthy_tool_fixer/main.py b/src/filthy_tool_fixer/main.py
index 46af30d..9e423b1 100644
--- a/src/filthy_tool_fixer/main.py
+++ b/src/filthy_tool_fixer/main.py
@@ -3,6 +3,7 @@
from __future__ import annotations
import asyncio
+import json
import time
from contextlib import asynccontextmanager
from typing import AsyncIterator
@@ -24,20 +25,43 @@
log = get_logger(__name__)
+
+class _RequestCounter:
+ """In-flight request counter with drain support for graceful shutdown.
+
+ All methods are synchronous — safe under asyncio's cooperative model
+ since no await can interleave between the read-modify-write steps.
+ """
+
+ def __init__(self) -> None:
+ self.count = 0
+ self.drain_event = asyncio.Event()
+ self.drain_event.set() # Initially drained
+
+ def enter(self) -> None:
+ self.count += 1
+ self.drain_event.clear()
+
+ def exit(self) -> None:
+ self.count -= 1
+ if self.count == 0:
+ self.drain_event.set()
+
+
# Module-level references set during lifespan
_orchestrator: ProxyOrchestrator | None = None
_profile_loader: ProfileLoader | None = None
_primary_backend: OllamaAdapter | None = None
_escalation_backend: OllamaAdapter | None = None
-_in_flight: int = 0
-_drain_event: asyncio.Event | None = None
+_requests = _RequestCounter()
_SHUTDOWN_DRAIN_TIMEOUT = 15.0 # Max seconds to wait for in-flight requests
+_STREAM_CHUNK_TIMEOUT = 60.0 # Max seconds to wait between stream chunks
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
- global _orchestrator, _profile_loader, _primary_backend, _escalation_backend, _drain_event
+ global _orchestrator, _profile_loader, _primary_backend, _escalation_backend
setup_logging()
log.info(
@@ -46,24 +70,27 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
escalation_url=settings.escalation_backend_url,
)
- _drain_event = asyncio.Event()
- _drain_event.set() # Initially "drained" (no requests)
-
# Load profiles
_profile_loader = ProfileLoader(settings.profiles_dir)
- # Initialize backends
+ # Initialize backends (tolerate startup failures for graceful degradation)
_primary_backend = OllamaAdapter(
base_url=settings.backend_url,
default_timeout=settings.backend_timeout,
)
- await _primary_backend.startup()
+ try:
+ await _primary_backend.startup()
+ except Exception:
+ log.exception("primary_backend_startup_failed")
_escalation_backend = OllamaAdapter(
base_url=settings.escalation_backend_url,
default_timeout=settings.escalation_timeout,
)
- await _escalation_backend.startup()
+ try:
+ await _escalation_backend.startup()
+ except Exception:
+ log.exception("escalation_backend_startup_failed")
_orchestrator = ProxyOrchestrator(
primary_backend=_primary_backend,
@@ -76,15 +103,24 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
yield
# Wait for in-flight requests to drain before closing backends
- if _in_flight > 0:
- log.info("shutdown_draining", in_flight=_in_flight)
+ if _requests.count > 0:
+ log.info("shutdown_draining", in_flight=_requests.count)
+ try:
+ await asyncio.wait_for(
+ _requests.drain_event.wait(), timeout=_SHUTDOWN_DRAIN_TIMEOUT
+ )
+ except asyncio.TimeoutError:
+ log.warning("shutdown_drain_timeout", in_flight=_requests.count)
+
+ # Shutdown backends with timeout to prevent hanging
+ for backend, name in [(_primary_backend, "primary"), (_escalation_backend, "escalation")]:
try:
- await asyncio.wait_for(_drain_event.wait(), timeout=_SHUTDOWN_DRAIN_TIMEOUT)
+ await asyncio.wait_for(backend.shutdown(), timeout=5.0)
except asyncio.TimeoutError:
- log.warning("shutdown_drain_timeout", in_flight=_in_flight)
+ log.warning(f"{name}_backend_shutdown_timeout")
+ except Exception:
+ log.exception(f"{name}_backend_shutdown_error")
- await _primary_backend.shutdown()
- await _escalation_backend.shutdown()
log.info("stopped")
@@ -135,8 +171,6 @@ async def list_models() -> JSONResponse:
@app.post("/v1/chat/completions", response_model=None)
async def chat_completions(request: Request):
- global _in_flight
-
if _orchestrator is None or _profile_loader is None:
return JSONResponse(
status_code=503,
@@ -155,8 +189,7 @@ async def chat_completions(request: Request):
status_code=413,
content={"error": {"message": "Request body too large", "type": "invalid_request_error"}},
)
- import json as _json
- body = _json.loads(body_bytes)
+ body = json.loads(body_bytes)
chat_request = ChatCompletionRequest.model_validate(body)
except Exception:
log.exception("invalid_request")
@@ -183,9 +216,7 @@ async def chat_completions(request: Request):
last_role=chat_request.messages[-1].role if chat_request.messages else "",
)
- _in_flight += 1
- if _drain_event is not None:
- _drain_event.clear()
+ _requests.enter()
try:
result = await asyncio.wait_for(
_orchestrator.handle_request(chat_request, profile),
@@ -207,9 +238,7 @@ async def chat_completions(request: Request):
content={"error": {"message": "Backend request failed", "type": "proxy_error"}},
)
finally:
- _in_flight -= 1
- if _in_flight == 0 and _drain_event is not None:
- _drain_event.set()
+ _requests.exit()
elapsed = time.monotonic() - start
@@ -217,7 +246,7 @@ async def chat_completions(request: Request):
if hasattr(result, "__aiter__"):
log.info("response_streaming", elapsed_ms=round(elapsed * 1000))
return StreamingResponse(
- result,
+ _timeout_stream(result),
media_type="text/event-stream",
headers={"X-FilthyToolFixer-Request-ID": rid},
)
@@ -254,19 +283,42 @@ async def chat_completions(request: Request):
)
+async def _timeout_stream(
+ stream: AsyncIterator[bytes],
+ chunk_timeout: float = _STREAM_CHUNK_TIMEOUT,
+) -> AsyncIterator[bytes]:
+ """Wrap a stream with per-chunk timeout to prevent hanging on stalled backends."""
+ aiter = stream.__aiter__()
+ while True:
+ try:
+ chunk = await asyncio.wait_for(aiter.__anext__(), timeout=chunk_timeout)
+ yield chunk
+ except StopAsyncIteration:
+ break
+ except asyncio.TimeoutError:
+ log.warning("stream_chunk_timeout", timeout=chunk_timeout)
+ error = {"error": {"message": "Backend stream stalled", "type": "proxy_timeout"}}
+ yield f"data: {json.dumps(error)}\n\n".encode()
+ yield b"data: [DONE]\n\n"
+ break
+ except asyncio.CancelledError:
+ raise
+
+
async def _synthesize_sse(response) -> AsyncIterator[bytes]:
"""Synthesize SSE events from a buffered ChatCompletionResponse.
Emits the response as a single chunk followed by [DONE].
"""
- import json
-
data = response.model_dump(exclude_none=True)
# Convert to streaming format
data["object"] = "chat.completion.chunk"
for choice in data.get("choices", []):
if "message" in choice:
delta = choice.pop("message")
+ # Ensure content field is present (some clients expect it even as null)
+ if "content" not in delta:
+ delta["content"] = None
# Streaming tool calls require an index field on each tool call
for i, tc in enumerate(delta.get("tool_calls", [])):
tc["index"] = i
diff --git a/src/filthy_tool_fixer/profiles/loader.py b/src/filthy_tool_fixer/profiles/loader.py
index ee92912..f2f1026 100644
--- a/src/filthy_tool_fixer/profiles/loader.py
+++ b/src/filthy_tool_fixer/profiles/loader.py
@@ -50,7 +50,7 @@ def _parse(self, data: dict) -> ModelProfile:
backend_url=model_data.get("backend_url", ""),
tool_calling=ToolCallingConfig(
system_suffix=tc_data.get("system_suffix", ""),
- temperature_override=tc_data.get("temperature_override", 0.0),
+ temperature_override=tc_data.get("temperature_override"),
tool_choice_override=tc_data.get("tool_choice_override", ""),
strip_thinking=tc_data.get("strip_thinking", False),
think_tag_pattern=tc_data.get("think_tag_pattern", ".*?"),
diff --git a/src/filthy_tool_fixer/profiles/types.py b/src/filthy_tool_fixer/profiles/types.py
index 9c57178..d12337b 100644
--- a/src/filthy_tool_fixer/profiles/types.py
+++ b/src/filthy_tool_fixer/profiles/types.py
@@ -16,7 +16,7 @@ class EscalationConfig:
@dataclass
class ToolCallingConfig:
system_suffix: str = ""
- temperature_override: float | None = 0.0
+ temperature_override: float | None = None # None = passthrough client temperature
strip_thinking: bool = False
think_tag_pattern: str = ".*?" # Regex pattern for thinking tags
tool_choice_override: str = "" # Override tool_choice sent to backend (e.g. "required")
diff --git a/src/filthy_tool_fixer/proxy.py b/src/filthy_tool_fixer/proxy.py
index 548f8f5..f140984 100644
--- a/src/filthy_tool_fixer/proxy.py
+++ b/src/filthy_tool_fixer/proxy.py
@@ -3,6 +3,7 @@
from __future__ import annotations
import asyncio
+import functools
import re
import time
from typing import AsyncIterator
@@ -22,8 +23,10 @@
log = get_logger(__name__)
-# Cache compiled think-tag patterns per profile
-_think_re_cache: dict[str, re.Pattern] = {}
+# Cache compiled think-tag patterns per profile (bounded)
+@functools.lru_cache(maxsize=32)
+def _compile_think_re(pattern: str) -> re.Pattern:
+ return re.compile(pattern, re.DOTALL)
# Regex patterns for condensing tool descriptions
_EXAMPLE_BLOCK_RE = re.compile(r".*?", re.DOTALL)
@@ -232,9 +235,7 @@ def _strip_thinking(self, response: ChatCompletionResponse, pattern: str) -> Cha
"""Remove thinking blocks from response content using the profile's pattern."""
if not pattern:
return response
- if pattern not in _think_re_cache:
- _think_re_cache[pattern] = re.compile(pattern, re.DOTALL)
- compiled = _think_re_cache[pattern]
+ compiled = _compile_think_re(pattern)
for choice in response.choices:
if choice.message and choice.message.content:
choice.message.content = compiled.sub("", choice.message.content).strip()
diff --git a/src/filthy_tool_fixer/retry/loop.py b/src/filthy_tool_fixer/retry/loop.py
index a6a1225..09105c9 100644
--- a/src/filthy_tool_fixer/retry/loop.py
+++ b/src/filthy_tool_fixer/retry/loop.py
@@ -7,6 +7,9 @@
import time
import uuid
from difflib import get_close_matches
+
+import httpx
+
from filthy_tool_fixer.backends.base import BackendAdapter
from filthy_tool_fixer.logging import get_logger
from filthy_tool_fixer.models import (
@@ -68,7 +71,7 @@ async def execute(
remaining = budget_remaining - elapsed
# Don't start a retry if insufficient time remains
- if attempt > 0 and remaining < 10.0:
+ if attempt > 0 and remaining < 20.0:
log.info("retry_budget_exhausted", attempt=attempt, remaining=remaining)
break
@@ -81,6 +84,23 @@ async def execute(
timeout=effective_timeout,
keep_alive=self._profile.tool_calling.keep_alive,
)
+ except httpx.HTTPStatusError as e:
+ status = e.response.status_code
+ if status < 500:
+ # 4xx = permanent error, don't retry
+ log.warning("backend_client_error", attempt=attempt, status=status)
+ break
+ # 5xx = transient, retry
+ log.warning("backend_server_error", attempt=attempt, status=status)
+ if attempt < max_retries:
+ continue
+ break
+ except (httpx.TimeoutException, httpx.ConnectError) as e:
+ # Transient network/timeout errors — retry
+ log.warning("backend_transient_error", attempt=attempt, error=str(e))
+ if attempt < max_retries:
+ continue
+ break
except Exception:
log.exception("backend_request_failed", attempt=attempt)
break
@@ -286,7 +306,8 @@ async def _try_escalation(
if failed_response:
# Add a brief summary of what went wrong, not the full failed conversation
summary = self._build_escalation_summary(failed_response, tools)
- messages.append(ChatMessage(role="user", content=summary))
+ if summary:
+ messages.append(ChatMessage(role="user", content=summary))
escalation_request = original_request.model_copy(
update={
@@ -379,7 +400,10 @@ def _extract_tool_calls(self, response: ChatCompletionResponse):
# Llama 4 pythonic: [func_name(param="val", param2=val2)]
# May contain multiple calls: [func1(a=1), func2(b=2), func3(c=3)]
_PYTHONIC_BRACKET_RE = re.compile(r"\[([^\]]+)\]", re.DOTALL)
- _PYTHONIC_CALL_RE = re.compile(r"(\w+)\(([^)]*)\)", re.DOTALL)
+ # Match function calls, allowing quoted strings to contain parens
+ _PYTHONIC_CALL_RE = re.compile(
+ r"""(\w+)\(((?:[^()"']*|"[^"]*"|'[^']*')*)\)""", re.DOTALL
+ )
def _rescue_tool_calls_from_text(
self,
@@ -411,8 +435,9 @@ def _rescue_tool_calls_from_text(
return rescued
# Try Llama 4 pythonic format: [func_name(param="val")]
- bracket_match = self._PYTHONIC_BRACKET_RE.search(content)
- if bracket_match:
+ # Try all bracket groups, not just the first (avoids false match on
+ # non-tool brackets like [1, 2, 3] appearing earlier in the text)
+ for bracket_match in self._PYTHONIC_BRACKET_RE.finditer(content):
inner = bracket_match.group(1)
pythonic_matches = self._PYTHONIC_CALL_RE.findall(inner)
rescued = self._parse_pythonic_calls(pythonic_matches, tool_names)
diff --git a/src/filthy_tool_fixer/validation/schema.py b/src/filthy_tool_fixer/validation/schema.py
index cd574cd..6678870 100644
--- a/src/filthy_tool_fixer/validation/schema.py
+++ b/src/filthy_tool_fixer/validation/schema.py
@@ -157,7 +157,7 @@ def _validate_single_call(
validator = jsonschema.Draft7Validator(schema)
for error in validator.iter_errors(args):
# Skip errors we already reported (required, additionalProperties)
- if error.validator in ("required",):
+ if error.validator in ("required", "additionalProperties"):
continue
errors.append(
ValidationError(