Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
10 changes: 8 additions & 2 deletions src/filthy_tool_fixer/backends/ollama.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
32 changes: 23 additions & 9 deletions src/filthy_tool_fixer/backends/vllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

from __future__ import annotations

import asyncio
import json
from typing import Any, AsyncIterator

Expand All @@ -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,
Expand Down Expand Up @@ -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

Expand All @@ -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

Expand All @@ -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:
Expand Down
108 changes: 80 additions & 28 deletions src/filthy_tool_fixer/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from __future__ import annotations

import asyncio
import json
import time
from contextlib import asynccontextmanager
from typing import AsyncIterator
Expand All @@ -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(
Expand All @@ -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,
Expand All @@ -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")


Expand Down Expand Up @@ -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,
Expand All @@ -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")
Expand All @@ -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),
Expand All @@ -207,17 +238,15 @@ 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

# Streaming response (non-tool passthrough)
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},
)
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/filthy_tool_fixer/profiles/loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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", "<think>.*?</think>"),
Expand Down
2 changes: 1 addition & 1 deletion src/filthy_tool_fixer/profiles/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = "<think>.*?</think>" # Regex pattern for thinking tags
tool_choice_override: str = "" # Override tool_choice sent to backend (e.g. "required")
Expand Down
11 changes: 6 additions & 5 deletions src/filthy_tool_fixer/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from __future__ import annotations

import asyncio
import functools
import re
import time
from typing import AsyncIterator
Expand All @@ -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"<example>.*?</example>", re.DOTALL)
Expand Down Expand Up @@ -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()
Expand Down
Loading