diff --git a/dev-notes/hugging-face-sticky-routing.md b/dev-notes/hugging-face-sticky-routing.md index 17cce4999..dc658a1f0 100644 --- a/dev-notes/hugging-face-sticky-routing.md +++ b/dev-notes/hugging-face-sticky-routing.md @@ -1,6 +1,7 @@ -# Hugging Face session-scoped provider routing +# Hugging Face session routing and cache adaptation -Issue: +Issues: , + ## What changed @@ -8,8 +9,9 @@ Tau pins a logical model from its built-in `huggingface` provider to one explicit Hugging Face Inference Provider. Users can configure a per-model `inference_providers` map in `~/.tau/providers.json`. Without a preference, the first request uses automatic routing; after it succeeds, Tau reads the -`x-inference-provider` response header and rebuilds the runtime with that explicit -suffix. The session index records the route for resume. +normalized `response_provider` value and rebuilds the runtime with that explicit +suffix. The session index records the route and whether it came from an explicit +choice or automatic routing. The OpenAI-compatible runtime now supports an internal logical-to-wire model alias. For example, the harness and persisted assistant messages continue to use @@ -18,11 +20,22 @@ For example, the harness and persisted assistant messages continue to use catalog and keeps context windows, capabilities, pricing, thinking controls, and model selection keyed by the logical model. -`/session` reports the pin. Route changes belong to the external Hugging Face -extension, which uses the public extension API. Switching logical models selects the new -model's configured pin or returns to automatic Hugging Face routing when none -exists. Existing records without the optional field remain compatible and pin -after their next successful automatic response. +Automatic routes now receive a bounded cache check. Requests below 4,096 prompt +tokens are ignored. The first eligible request on a route is a cold warm-up; two +later append-only requests with absent cache telemetry or an explicit zero cache +read trigger candidate discovery. A positive cache read retains the route. Tau +tries at most three routes across at most nine eligible requests, and keeps the +current route if discovery or all remaining candidates fail. +Candidate discovery filters the model API mapping to validated `status: live`, +`task: conversational` suffixes, sorts them lexicographically, and skips routes +already attempted in the session. + +`/session` reports the pin and evaluation phase. Every automatic route change is +a typed coding-session event consumed by print and TUI frontends. Manual route +selection remains available through the external Hugging Face extension. An +explicit route locks evaluation; `/route automatic` resets it. Historical +records with a pin but no source field are treated as explicit, so an upgrade +cannot silently take over a user's existing route. ## Why it exists @@ -32,17 +45,26 @@ full misses only seconds after large cache reads, followed immediately by anothe large hit. That pattern is consistent with requests moving between provider or worker cache domains rather than normal TTL expiry. +In the motivating GLM-5.2 comparison, `deepinfra` reported roughly 99% reuse +after its cold request. Five append-only requests through `scaleway` returned +null cache details for an approximately 12.5k-token prefix and incurred about +$0.13 in Hugging Face billing, consistent with processing the prompt as fresh. +That is evidence about reported reuse and billing, not proof that a backend has +no internal cache. + An explicit `:` suffix narrows one source of routing changes. It does not guarantee a cache hit: the selected provider can still evict entries, load-balance across workers, or expire them. ## Architecture -Provider preferences, session metadata, and model-route selection remain in -`tau_coding`. The reusable `tau_agent` harness receives the logical model and has -no Hugging Face-specific behavior. `tau_ai` only gains a provider-neutral -`model_aliases` transport option, used to put a different model ID in the wire -payload while preserving the logical model in normalized events. +Provider preferences, adaptive policy, session metadata, and model-route +selection remain in `tau_coding`. State snapshots use the existing session +`CustomEntry` mechanism, so evidence, unavailable candidates, and transitions +resume with the active branch. The reusable `tau_agent` harness receives the +logical model and has no Hugging Face-specific behavior. `tau_ai` and the +provider-neutral `Usage` model expose only whether a cache-read counter was +reported, which distinguishes absent telemetry from a reported zero. [Hugging Face's own Chat UI](https://github.com/huggingface/chat-ui/blob/main/src/lib/server/endpoints/openai/endpointOai.ts) consumes `x-inference-provider` from OpenAI-compatible responses, so Tau uses @@ -50,12 +72,11 @@ that header rather than guessing which mapping is fastest. The pin is committed only after a successful stream. Existing OpenAI-compatible retries keep the same wire model and stop retrying after streamed model output. -The first version deliberately does not automatically fail over a stale explicit -pin or send provider-specific cache-affinity fields. Safe failover also needs a -user-visible reroute event plus durable retry/reroute telemetry; silently falling -back would hide temporary cache-locality loss. Users can explicitly reset with -the Hugging Face extension. Backing providers differ in accepted affinity fields, so no -unknown field is enabled for the entire gateway. +Tau deliberately does not fail over an explicit pin or send provider-specific +cache-affinity fields. Errors, aborts, model changes, compaction boundaries, +route mismatches, and non-append-only contexts do not add cache-failure evidence. +Backing providers differ in accepted affinity fields, so no unknown field is +enabled for the entire gateway. ## Configure and validate diff --git a/src/tau_agent/messages.py b/src/tau_agent/messages.py index efe7089fe..680abb5ec 100644 --- a/src/tau_agent/messages.py +++ b/src/tau_agent/messages.py @@ -48,6 +48,7 @@ class Usage(WireModel): input: int = 0 output: int = 0 cache_read: int = 0 + cache_read_reported: bool | None = None cache_write: int = 0 cache_write_1h: int | None = None reasoning: int | None = None diff --git a/src/tau_ai/openai_compatible.py b/src/tau_ai/openai_compatible.py index b8b3930b4..648ffa489 100644 --- a/src/tau_ai/openai_compatible.py +++ b/src/tau_ai/openai_compatible.py @@ -1286,6 +1286,7 @@ def _parse_chunk_usage(raw: Mapping[str, Any]) -> Usage: # 0 does not fall through. if cached_tokens is None: cached_tokens = _int_or_none(raw.get("prompt_cache_hit_tokens")) + cache_read_reported = cached_tokens is not None cache_read = cached_tokens or 0 fresh_input = max(0, prompt_tokens - cache_read - cache_write) output = _int_or_zero(raw.get("completion_tokens")) @@ -1297,6 +1298,7 @@ def _parse_chunk_usage(raw: Mapping[str, Any]) -> Usage: input=fresh_input, output=output, cache_read=cache_read, + cache_read_reported=cache_read_reported, cache_write=cache_write, reasoning=reasoning, total_tokens=fresh_input + output + cache_read + cache_write, @@ -1322,6 +1324,11 @@ def _usage_from_responses_event(chunk: Mapping[str, Any]) -> Usage | None: if isinstance(input_details, Mapping) else 0 ) + cache_read_reported = ( + _int_or_none(input_details.get("cached_tokens")) is not None + if isinstance(input_details, Mapping) + else False + ) cache_write = ( _int_or_zero(input_details.get("cache_write_tokens")) if isinstance(input_details, Mapping) @@ -1339,6 +1346,7 @@ def _usage_from_responses_event(chunk: Mapping[str, Any]) -> Usage | None: input=max(0, _int_or_zero(raw.get("input_tokens")) - cache_read - cache_write), output=_int_or_zero(raw.get("output_tokens")), cache_read=cache_read, + cache_read_reported=cache_read_reported, cache_write=cache_write, reasoning=reasoning, total_tokens=_int_or_zero(raw.get("total_tokens")), diff --git a/src/tau_coding/cli.py b/src/tau_coding/cli.py index cd0d08828..8f3140074 100644 --- a/src/tau_coding/cli.py +++ b/src/tau_coding/cli.py @@ -54,7 +54,12 @@ export_session_artifact, normalize_export_format, ) -from tau_coding.session_manager import CodingSessionRecord, SessionManager, validate_session_id +from tau_coding.session_manager import ( + CodingSessionRecord, + HuggingFaceRouteMode, + SessionManager, + validate_session_id, +) from tau_coding.shell_config import load_shell_settings from tau_coding.tui import run_tui_app from tau_coding.update_check import ( @@ -912,12 +917,15 @@ async def run_openai_print_mode( provider_name=provider_name if explicit_selection else record.provider_name, model=model if explicit_selection else record.model, ) - inference_provider = ( - record.inference_provider - if resume_session_id is not None + uses_resumed_route = ( + resume_session_id is not None and record.provider_name == "huggingface" and selection.provider.name == "huggingface" and record.model == selection.model + ) + inference_provider = ( + record.inference_provider + if uses_resumed_route else selection.provider.inference_providers.get(selection.model) if isinstance(selection.provider, OpenAICompatibleProviderConfig) and selection.provider.name == "huggingface" @@ -941,6 +949,9 @@ async def run_openai_print_mode( session_manager=manager, provider_name=selection.provider.name, inference_provider=inference_provider, + inference_provider_mode=( + record.inference_provider_mode if uses_resumed_route else None + ), provider_settings=settings, runtime_provider_config=selection.provider, shell_command_prefix=shell_settings.shell_command_prefix, @@ -1023,6 +1034,7 @@ async def run_print_mode( session_manager: SessionManager | None = None, provider_name: str = DEFAULT_PROVIDER_NAME, inference_provider: str | None = None, + inference_provider_mode: HuggingFaceRouteMode | None = None, provider_settings: ProviderSettings | None = None, runtime_provider_config: ProviderConfig | None = None, shell_command_prefix: str | None = None, @@ -1051,6 +1063,7 @@ async def run_print_mode( session_manager=session_manager, provider_name=provider_name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, provider_settings=provider_settings, runtime_provider_config=runtime_provider_config, shell_command_prefix=shell_command_prefix, diff --git a/src/tau_coding/commands.py b/src/tau_coding/commands.py index d84847c54..3af1e9ca0 100644 --- a/src/tau_coding/commands.py +++ b/src/tau_coding/commands.py @@ -436,6 +436,9 @@ def _status_command(context: CommandContext) -> CommandResult: if session.provider_name == "huggingface": route = getattr(session, "inference_provider", None) or "automatic" lines.append(f"Hugging Face inference provider: {route}") + routing_status = getattr(session, "huggingface_routing_status", None) + if routing_status: + lines.append(f"Hugging Face cache routing: {routing_status}") context_window_source = getattr(session, "context_window_source", None) if context_window_source: lines.append(f"Context window source: {context_window_source}") diff --git a/src/tau_coding/events.py b/src/tau_coding/events.py index 3216dd470..e41b81884 100644 --- a/src/tau_coding/events.py +++ b/src/tau_coding/events.py @@ -74,6 +74,22 @@ class AutoRetryEndEvent(WireModel): final_error: str | None = Field(None) +class HuggingFaceRouteEvent(WireModel): + type: Literal["huggingface_route"] = "huggingface_route" + status: Literal["changed", "exhausted"] + previous_route: str | None = None + route: str + reason: str + + @property + def display_text(self) -> str: + """Return the concise route notice shown by human frontends.""" + if self.status == "changed": + previous = self.previous_route or "automatic" + return f"Hugging Face route changed: {previous} -> {self.route} ({self.reason})" + return f"Hugging Face route evaluation stopped on {self.route}: {self.reason}" + + type SessionOwnEvent = Annotated[ SessionAgentEndEvent | AgentSettledEvent @@ -84,7 +100,8 @@ class AutoRetryEndEvent(WireModel): | SessionInfoChangedEvent | ThinkingLevelChangedEvent | AutoRetryStartEvent - | AutoRetryEndEvent, + | AutoRetryEndEvent + | HuggingFaceRouteEvent, Field(discriminator="type"), ] type CodingSessionEvent = AgentEvent | SessionOwnEvent diff --git a/src/tau_coding/huggingface_routing.py b/src/tau_coding/huggingface_routing.py new file mode 100644 index 000000000..872cb1943 --- /dev/null +++ b/src/tau_coding/huggingface_routing.py @@ -0,0 +1,367 @@ +"""Bounded cache-aware routing policy for Tau's built-in Hugging Face provider.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from hashlib import sha256 +from json import dumps +from typing import Literal, cast +from urllib.parse import quote + +import httpx +from pydantic import BaseModel, ConfigDict, Field, ValidationError + +from tau_agent.messages import AgentMessage, Usage +from tau_agent.tools import AgentTool +from tau_agent.types import JSONValue +from tau_coding.provider_config import ( + ProviderConfigError, + validate_huggingface_inference_provider, +) + +HF_CACHE_ROUTING_NAMESPACE = "tau.huggingface-cache-routing" +HF_CACHE_MIN_PROMPT_TOKENS = 4_096 +HF_CACHE_PROBES_PER_ROUTE = 2 +HF_CACHE_MAX_CANDIDATES = 3 +HF_CACHE_MAX_ELIGIBLE_REQUESTS = 9 +HF_MODEL_API_BASE_URL = "https://huggingface.co/api/models" + +type HuggingFaceRoutingPhase = Literal["evaluating", "retained", "reroute", "exhausted"] + + +class RequestContextFingerprint(BaseModel): + """Minimal durable evidence for append-only request comparability.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + static_digest: str + message_count: int = Field(ge=0) + message_digest: str + + +class HuggingFaceRoutingState(BaseModel): + """Versioned session-owned state for Hugging Face adaptive routing.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + version: Literal[1] = 1 + model: str + route: str | None = None + phase: HuggingFaceRoutingPhase = "evaluating" + attempted_routes: tuple[str, ...] = () + unavailable_routes: tuple[str, ...] = () + eligible_requests: int = Field(default=0, ge=0) + absent_probes: int = Field(default=0, ge=0) + zero_probes: int = Field(default=0, ge=0) + last_context: RequestContextFingerprint | None = None + last_reason: str | None = None + + @classmethod + def automatic(cls, model: str, *, route: str | None = None) -> HuggingFaceRoutingState: + """Create a fresh automatic-routing evaluator for one logical model.""" + attempted = (route,) if route is not None else () + return cls(model=model, route=route, attempted_routes=attempted) + + def to_custom_data(self) -> dict[str, JSONValue]: + """Return a JSON-compatible snapshot for an application custom entry.""" + return cast(dict[str, JSONValue], self.model_dump(mode="json")) + + @classmethod + def from_custom_data(cls, data: Mapping[str, JSONValue]) -> HuggingFaceRoutingState | None: + """Load a supported snapshot, ignoring malformed or future versions.""" + if data.get("version") != 1: + return None + try: + return cls.model_validate(data) + except ValidationError: + return None + + +def observe_huggingface_cache( + state: HuggingFaceRoutingState, + usage: Usage, + *, + system: str, + messages: Sequence[AgentMessage], + tools: Sequence[AgentTool], +) -> HuggingFaceRoutingState: + """Apply one successful response's cache evidence to an automatic state.""" + if state.phase != "evaluating": + return state + if state.route is None: + return state.model_copy(update={"last_reason": "waiting for automatic route resolution"}) + + prompt_tokens = usage.input + usage.cache_read + usage.cache_write + if prompt_tokens < HF_CACHE_MIN_PROMPT_TOKENS: + return state.model_copy( + update={ + "last_reason": ( + f"prompt below {HF_CACHE_MIN_PROMPT_TOKENS}-token eligibility threshold" + ) + } + ) + + context = _request_context_fingerprint(system=system, messages=messages, tools=tools) + if state.last_context is None: + eligible_requests = state.eligible_requests + 1 + cold_update: dict[str, object] = { + "eligible_requests": eligible_requests, + "last_context": context, + "last_reason": "cold route warm-up", + } + if usage.cache_read_reported is True and usage.cache_read > 0: + cold_update.update(phase="retained", last_reason="positive cache reuse reported") + else: + cold_update.update( + _budget_exhaustion_update(state, eligible_requests=eligible_requests) + ) + return state.model_copy(update=cold_update) + + if not _context_extends( + state.last_context, + system=system, + messages=messages, + tools=tools, + ): + return state.model_copy( + update={ + "absent_probes": 0, + "zero_probes": 0, + "last_context": context, + "last_reason": "request context was not append-only comparable", + } + ) + + eligible_requests = state.eligible_requests + 1 + if usage.cache_read_reported is True and usage.cache_read > 0: + return state.model_copy( + update={ + "phase": "retained", + "eligible_requests": eligible_requests, + "last_context": context, + "last_reason": "positive cache reuse reported", + } + ) + + absent_probes = state.absent_probes + (usage.cache_read_reported is not True) + zero_probes = state.zero_probes + (usage.cache_read_reported is True) + failed_probes = absent_probes + zero_probes + update: dict[str, object] = { + "eligible_requests": eligible_requests, + "absent_probes": absent_probes, + "zero_probes": zero_probes, + "last_context": context, + "last_reason": ( + "no effective reported cache reuse after warmed probes " + f"(telemetry absent: {absent_probes}, reported zero: {zero_probes})" + ), + } + budget_update = _budget_exhaustion_update(state, eligible_requests=eligible_requests) + if budget_update: + update.update(budget_update) + elif failed_probes >= HF_CACHE_PROBES_PER_ROUTE: + update["phase"] = "reroute" + return state.model_copy(update=update) + + +def resolve_automatic_huggingface_route( + state: HuggingFaceRoutingState, + route: str, +) -> HuggingFaceRoutingState: + """Record the first route resolved by Hugging Face automatic routing.""" + normalized = validate_huggingface_inference_provider(route) + attempted = _append_unique(state.attempted_routes, normalized) + return state.model_copy( + update={ + "route": normalized, + "phase": "evaluating", + "attempted_routes": attempted, + "absent_probes": 0, + "zero_probes": 0, + "last_context": None, + "last_reason": "automatic route resolved", + } + ) + + +def reroute_huggingface_state( + state: HuggingFaceRoutingState, + route: str, + *, + reason: str, +) -> HuggingFaceRoutingState: + """Start a cold bounded probe on a replacement route.""" + normalized = validate_huggingface_inference_provider(route) + return state.model_copy( + update={ + "route": normalized, + "phase": "evaluating", + "attempted_routes": _append_unique(state.attempted_routes, normalized), + "absent_probes": 0, + "zero_probes": 0, + "last_context": None, + "last_reason": reason, + } + ) + + +def mark_huggingface_route_unavailable( + state: HuggingFaceRoutingState, + route: str, + *, + reason: str | None = None, +) -> HuggingFaceRoutingState: + """Consume one candidate slot without counting it as cache evidence.""" + normalized = validate_huggingface_inference_provider(route) + update: dict[str, object] = { + "attempted_routes": _append_unique(state.attempted_routes, normalized), + "unavailable_routes": _append_unique(state.unavailable_routes, normalized), + } + if reason is not None: + update["last_reason"] = reason + if normalized == state.route: + update["phase"] = "reroute" + return state.model_copy(update=update) + + +def exhaust_huggingface_routing( + state: HuggingFaceRoutingState, + *, + reason: str, +) -> HuggingFaceRoutingState: + """Stop automatic route changes while leaving the current route usable.""" + return state.model_copy(update={"phase": "exhausted", "last_reason": reason}) + + +def next_huggingface_routes( + state: HuggingFaceRoutingState, + discovered_routes: Sequence[str], +) -> tuple[str, ...]: + """Return deterministic untried routes within the total candidate budget.""" + remaining = max(0, HF_CACHE_MAX_CANDIDATES - len(state.attempted_routes)) + if remaining == 0: + return () + attempted = set(state.attempted_routes) + candidates: set[str] = set() + for route in discovered_routes: + try: + normalized = validate_huggingface_inference_provider(route) + except ProviderConfigError: + continue + if normalized not in attempted: + candidates.add(normalized) + return tuple(sorted(candidates))[:remaining] + + +async def discover_huggingface_routes( + model: str, + *, + client: httpx.AsyncClient | None = None, +) -> tuple[str, ...]: + """Discover live conversational provider suffixes for one logical model.""" + url = f"{HF_MODEL_API_BASE_URL}/{quote(model, safe='/')}" + owns_client = client is None + active_client = client or httpx.AsyncClient(timeout=10.0) + try: + response = await active_client.get( + url, + params={"expand[]": "inferenceProviderMapping"}, + timeout=10.0, + ) + response.raise_for_status() + payload = response.json() + finally: + if owns_client: + await active_client.aclose() + + if not isinstance(payload, Mapping): + return () + raw_mapping = payload.get("inferenceProviderMapping") + if not isinstance(raw_mapping, Mapping): + return () + routes: set[str] = set() + for raw_route, raw_details in raw_mapping.items(): + if not isinstance(raw_route, str) or not isinstance(raw_details, Mapping): + continue + if raw_details.get("status") != "live" or raw_details.get("task") != "conversational": + continue + try: + routes.add(validate_huggingface_inference_provider(raw_route)) + except ProviderConfigError: + continue + return tuple(sorted(routes)) + + +def _request_context_fingerprint( + *, + system: str, + messages: Sequence[AgentMessage], + tools: Sequence[AgentTool], +) -> RequestContextFingerprint: + static_payload = { + "system": system, + "tools": [ + { + "name": tool.name, + "description": tool.description, + "parameters": dict(tool.parameters), + } + for tool in tools + ], + } + return RequestContextFingerprint( + static_digest=_json_digest(static_payload), + message_count=len(messages), + message_digest=_message_digest(messages), + ) + + +def _context_extends( + previous: RequestContextFingerprint, + *, + system: str, + messages: Sequence[AgentMessage], + tools: Sequence[AgentTool], +) -> bool: + current = _request_context_fingerprint(system=system, messages=messages, tools=tools) + return ( + current.static_digest == previous.static_digest + and current.message_count > previous.message_count + and _message_digest(messages[: previous.message_count]) == previous.message_digest + ) + + +def _message_digest(messages: Sequence[AgentMessage]) -> str: + return _json_digest( + [message.model_dump(mode="json", by_alias=True, exclude_none=False) for message in messages] + ) + + +def _json_digest(value: object) -> str: + encoded = dumps( + value, + ensure_ascii=True, + sort_keys=True, + separators=(",", ":"), + ).encode("utf-8") + return sha256(encoded).hexdigest() + + +def _append_unique(values: tuple[str, ...], value: str) -> tuple[str, ...]: + return values if value in values else (*values, value) + + +def _budget_exhaustion_update( + state: HuggingFaceRoutingState, + *, + eligible_requests: int, +) -> dict[str, object]: + if len(state.attempted_routes) >= HF_CACHE_MAX_CANDIDATES: + if eligible_requests >= HF_CACHE_MAX_ELIGIBLE_REQUESTS: + return { + "phase": "exhausted", + "last_reason": "candidate budget and eligible request budget exhausted", + } + elif eligible_requests >= HF_CACHE_MAX_ELIGIBLE_REQUESTS: + return {"phase": "exhausted", "last_reason": "eligible request budget exhausted"} + return {} diff --git a/src/tau_coding/rendering/plain.py b/src/tau_coding/rendering/plain.py index 7a9f7f338..5b7914588 100644 --- a/src/tau_coding/rendering/plain.py +++ b/src/tau_coding/rendering/plain.py @@ -4,7 +4,7 @@ from tau_agent.events import MessageEndEvent from tau_agent.messages import AssistantMessage -from tau_coding.events import CodingSessionEvent +from tau_coding.events import CodingSessionEvent, HuggingFaceRouteEvent class FinalTextRenderer: @@ -14,6 +14,9 @@ def __init__(self) -> None: self._error_messages: list[str] = [] def render(self, event: CodingSessionEvent) -> None: + if isinstance(event, HuggingFaceRouteEvent): + typer.echo(event.display_text, err=True) + return if not isinstance(event, MessageEndEvent) or not isinstance( event.message, AssistantMessage ): diff --git a/src/tau_coding/rendering/transcript.py b/src/tau_coding/rendering/transcript.py index ea22aa9e0..e8042a656 100644 --- a/src/tau_coding/rendering/transcript.py +++ b/src/tau_coding/rendering/transcript.py @@ -14,7 +14,7 @@ ) from tau_agent.messages import AssistantMessage, CustomMessage, ToolCall from tau_ai.events import TextDeltaEvent -from tau_coding.events import AutoRetryStartEvent, CodingSessionEvent +from tau_coding.events import AutoRetryStartEvent, CodingSessionEvent, HuggingFaceRouteEvent from tau_coding.extensions.api import CustomMessageMarkup from tau_coding.tui.state import format_tool_call_block @@ -53,6 +53,10 @@ def render(self, event: CodingSessionEvent) -> None: self._newline() self._console.print(Text(f"… {event.error_message}", style="bright_black")) return + if isinstance(event, HuggingFaceRouteEvent): + self._newline() + self._console.print(Text(event.display_text, style="bright_black")) + return if isinstance(event, ToolExecutionEndEvent): status = "✗" if event.is_error else "✓" style = "red" if event.is_error else "green" diff --git a/src/tau_coding/session.py b/src/tau_coding/session.py index 6cacbcde1..7a9c2d8d5 100644 --- a/src/tau_coding/session.py +++ b/src/tau_coding/session.py @@ -3,7 +3,7 @@ from __future__ import annotations import string -from collections.abc import AsyncIterator, Callable, Mapping +from collections.abc import AsyncIterator, Callable from contextlib import suppress from dataclasses import dataclass, replace from pathlib import Path @@ -68,10 +68,24 @@ CodingSessionEvent, CompactionEndEvent, CompactionStartEvent, + HuggingFaceRouteEvent, QueueUpdateEvent, SessionAgentEndEvent, ) from tau_coding.extensions.runtime import ExtensionRuntime +from tau_coding.huggingface_routing import ( + HF_CACHE_MAX_ELIGIBLE_REQUESTS, + HF_CACHE_PROBES_PER_ROUTE, + HF_CACHE_ROUTING_NAMESPACE, + HuggingFaceRoutingState, + discover_huggingface_routes, + exhaust_huggingface_routing, + mark_huggingface_route_unavailable, + next_huggingface_routes, + observe_huggingface_cache, + reroute_huggingface_state, + resolve_automatic_huggingface_route, +) from tau_coding.paths import TauPaths from tau_coding.project_trust import ( CanonicalProjectPath, @@ -121,7 +135,7 @@ export_session_artifact, normalize_export_format, ) -from tau_coding.session_manager import SessionManager +from tau_coding.session_manager import HuggingFaceRouteMode, SessionManager from tau_coding.session_stats import SessionStats, calculate_session_stats from tau_coding.skills import Skill, expand_skill_command, load_skills_with_diagnostics from tau_coding.system_prompt import ( @@ -252,6 +266,7 @@ class CodingSessionConfig: command_registry: CommandRegistry | None = None provider_name: str = "openai" inference_provider: str | None = None + inference_provider_mode: HuggingFaceRouteMode | None = None provider_settings: ProviderSettings | None = None runtime_provider_config: ProviderConfig | None = None auto_compact_token_threshold: int | None = None @@ -331,6 +346,9 @@ def __init__( self._command_registry = command_registry or create_default_command_registry() self._provider_name = config.provider_name self._inference_provider = config.inference_provider + self._inference_provider_mode: HuggingFaceRouteMode | None = None + self._huggingface_routing_state: HuggingFaceRoutingState | None = None + self._restore_huggingface_routing(config, state) self._provider_settings = config.provider_settings self._runtime_provider_config = config.runtime_provider_config self._resource_paths = resource_paths_with_cwd(config.resource_paths, config.cwd) @@ -560,6 +578,125 @@ def inference_provider(self) -> str | None: """Return the pinned Hugging Face backing provider, if any.""" return self._inference_provider + @property + def huggingface_routing_status(self) -> str | None: + """Return the concise adaptive-routing state shown by `/session`.""" + if self._provider_name != "huggingface": + return None + if self._inference_provider_mode == "explicit": + return "explicit route lock" + state = self._huggingface_routing_state + if state is None: + return None + if state.phase == "retained": + return "automatic; retained after positive cache reuse" + if state.phase == "exhausted": + return f"automatic; evaluation stopped ({state.last_reason or 'budget exhausted'})" + if state.route is None: + return "automatic; waiting for route resolution" + if state.phase == "reroute": + return "automatic; replacement route pending" + failed_probes = state.absent_probes + state.zero_probes + return ( + f"automatic; evaluating {state.route} " + f"({failed_probes}/{HF_CACHE_PROBES_PER_ROUTE} warmed misses, " + f"{state.eligible_requests}/{HF_CACHE_MAX_ELIGIBLE_REQUESTS} eligible requests)" + ) + + def _restore_huggingface_routing( + self, + config: CodingSessionConfig, + state: SessionState, + ) -> None: + if self._provider_name != "huggingface": + return + mode = config.inference_provider_mode + if mode is None: + mode = "explicit" if config.inference_provider is not None else "automatic" + if mode == "explicit": + routing_state = None + else: + # A missing stored mode is also the durable marker for a fresh automatic reset. + latest_state = ( + None + if config.inference_provider_mode is None + else next( + ( + restored + for entry in reversed(state.custom_entries) + if entry.namespace == HF_CACHE_ROUTING_NAMESPACE + and (restored := HuggingFaceRoutingState.from_custom_data(entry.data)) + is not None + and restored.model == self.model + ), + None, + ) + ) + routing_state = ( + latest_state + if latest_state is not None and latest_state.route == config.inference_provider + else HuggingFaceRoutingState.automatic( + self.model, + route=config.inference_provider, + ) + ) + self._inference_provider = routing_state.route + self._inference_provider_mode = mode + self._huggingface_routing_state = routing_state + self._config = replace( + config, + inference_provider=self._inference_provider, + inference_provider_mode=mode, + ) + + def _reset_huggingface_routing( + self, + route: str | None, + *, + mode: HuggingFaceRouteMode | None = None, + ) -> None: + if self._provider_name != "huggingface": + self._inference_provider = None + self._inference_provider_mode = None + self._huggingface_routing_state = None + else: + resolved_mode = mode or ("explicit" if route is not None else "automatic") + self._inference_provider = route + self._inference_provider_mode = resolved_mode + self._huggingface_routing_state = ( + None + if resolved_mode == "explicit" + else HuggingFaceRoutingState.automatic(self.model, route=route) + ) + self._config = replace( + self._config, + inference_provider=self._inference_provider, + inference_provider_mode=self._inference_provider_mode, + ) + + def _indexed_inference_provider_mode(self) -> HuggingFaceRouteMode | None: + state = self._huggingface_routing_state + if ( + self._provider_name == "huggingface" + and self._inference_provider_mode == "automatic" + and self._inference_provider is None + and state == HuggingFaceRoutingState.automatic(self.model) + ): + return None + return self._inference_provider_mode + + def _touch_session_index(self) -> None: + if self._config.session_id is None or self._config.session_manager is None: + return + self._config.session_manager.touch_session( + self._config.session_id, + model=self.model, + provider_name=self.provider_name, + inference_provider=self._inference_provider, + inference_provider_mode=self._indexed_inference_provider_mode(), + preserve_inference_provider=False, + ) + @property def available_providers(self) -> tuple[str, ...]: """Return provider names Tau can call with available credentials.""" @@ -706,6 +843,8 @@ async def branch_to_entry( await self._refresh_persisted_state(leaf_id=target_id) history_repair = await self._persist_active_tool_history_repairs() + if self._huggingface_routing_state is not None: + await self._persist_huggingface_routing_state(self._huggingface_routing_state) if history_repair is None: self._harness.replace_messages(self._state.messages) self._invalidate_context_usage_cache() @@ -1095,19 +1234,12 @@ def set_model(self, model: str) -> None: if provider is not None: validate_provider_model(provider, model) self._harness.config.model = model - self._inference_provider = _configured_inference_provider(provider, model) + self._reset_huggingface_routing(_configured_inference_provider(provider, model)) self._sync_thinking_level_to_active_model() self._refresh_runtime_provider() self._sync_image_support() self._persist_default_model_choice() - if self._config.session_id is not None and self._config.session_manager is not None: - self._config.session_manager.touch_session( - self._config.session_id, - model=model, - provider_name=self.provider_name, - inference_provider=self._inference_provider, - preserve_inference_provider=False, - ) + self._touch_session_index() async def apply_startup_model_override(self, model: str) -> None: """Activate and persist an explicit startup model before the next turn.""" @@ -1118,7 +1250,7 @@ async def apply_startup_model_override(self, model: str) -> None: return self._harness.config.model = model - self._inference_provider = _configured_inference_provider(provider, model) + self._reset_huggingface_routing(_configured_inference_provider(provider, model)) self._sync_thinking_level_to_active_model() self._refresh_runtime_provider() self._sync_image_support() @@ -1140,6 +1272,7 @@ def set_inference_provider(self, route: str | None) -> str: "Inference-provider routing requires the huggingface provider" ) normalized = validate_huggingface_inference_provider(route) if route is not None else None + mode: HuggingFaceRouteMode = "explicit" if normalized is not None else "automatic" provider, provider_config = self._build_runtime_provider( inference_provider=normalized, ) @@ -1150,10 +1283,10 @@ def set_inference_provider(self, route: str | None) -> str: model=self.model, provider_name=self.provider_name, inference_provider=normalized, + inference_provider_mode=None if mode == "automatic" else mode, preserve_inference_provider=False, ) - self._inference_provider = normalized - self._config = replace(self._config, inference_provider=normalized) + self._reset_huggingface_routing(normalized, mode=mode) self._activate_runtime_provider(provider, provider_config) return normalized or "automatic (will pin after the next successful response)" @@ -1239,33 +1372,21 @@ def _set_provider_model( model=model, thinking_level=thinking_level, inference_provider=_configured_inference_provider(provider_config, model), - response_headers_observer=( - self._observe_response_headers - if provider_config.name == "huggingface" - else None - ), ) except RuntimeError as exc: raise ProviderConfigError(str(exc)) from exc self._owned_providers.append(provider) self._harness.config.provider = provider self._provider_name = provider_config.name - self._inference_provider = _configured_inference_provider(provider_config, model) self._runtime_provider_config = provider_config self._invalidate_runtime_model_limits() self._harness.config.model = model + self._reset_huggingface_routing(_configured_inference_provider(provider_config, model)) self._thinking_level = thinking_level self._sync_image_support() if persist_default: self._persist_default_model_choice() - if self._config.session_id is not None and self._config.session_manager is not None: - self._config.session_manager.touch_session( - self._config.session_id, - model=model, - provider_name=self.provider_name, - inference_provider=self._inference_provider, - preserve_inference_provider=False, - ) + self._touch_session_index() async def set_thinking_level(self, level: str) -> str: """Persist and activate a thinking mode for future turns.""" @@ -1368,37 +1489,6 @@ def _persist_thinking_level_choice(self) -> None: except ProviderConfigError: return - def _observe_response_headers(self, headers: Mapping[str, str]) -> None: - if self.provider_name != "huggingface" or self._inference_provider is not None: - return - route = next( - (value for key, value in headers.items() if key.casefold() == "x-inference-provider"), - None, - ) - if route is None: - return - try: - route = validate_huggingface_inference_provider(route) - except ProviderConfigError: - return - provider, provider_config = self._build_runtime_provider( - inference_provider=route, - ) - # Track staged providers immediately so a later index-write failure does - # not leak a provider-owned client. The active runtime remains unchanged. - self._owned_providers.append(provider) - if self._config.session_manager is not None and self._config.session_id is not None: - self._config.session_manager.touch_session( - self._config.session_id, - model=self.model, - provider_name=self.provider_name, - inference_provider=route, - preserve_inference_provider=False, - ) - self._inference_provider = route - self._config = replace(self._config, inference_provider=route) - self._activate_runtime_provider(provider, provider_config) - def _build_runtime_provider( self, *, @@ -1415,11 +1505,6 @@ def _build_runtime_provider( model=self.model, thinking_level=self._thinking_level, inference_provider=inference_provider, - response_headers_observer=( - self._observe_response_headers - if provider_config.name == "huggingface" - else None - ), ) except RuntimeError as exc: raise ProviderConfigError(str(exc)) from exc @@ -1443,6 +1528,271 @@ def _refresh_runtime_provider(self) -> None: self._owned_providers.append(provider) self._activate_runtime_provider(provider, provider_config) + async def _observe_huggingface_response( + self, + message: AssistantMessage, + ) -> HuggingFaceRouteEvent | None: + state = self._huggingface_routing_state + if ( + state is None + or self._inference_provider_mode != "automatic" + or state.model != self.model + or message.stop_reason == "aborted" + ): + return None + + response_route: str | None = None + if message.response_provider is not None: + try: + response_route = validate_huggingface_inference_provider(message.response_provider) + except ProviderConfigError: + return None + if state.route is not None and response_route not in {None, state.route}: + return None + + if message.stop_reason == "error": + if ( + len(state.attempted_routes) > 1 + and state.route is not None + and not message.content + and not is_context_overflow_error(message) + ): + failed = mark_huggingface_route_unavailable( + state, + state.route, + reason=f"{state.route} failed before producing output", + ) + await self._persist_huggingface_routing_state(failed, expected_state=state) + return None + + if state.route is None: + if response_route is None: + return None + try: + provider, provider_config = self._build_runtime_provider( + inference_provider=response_route, + ) + except ProviderConfigError: + reason = f"{response_route} could not be pinned locally" + unavailable = mark_huggingface_route_unavailable( + state, + response_route, + reason=reason, + ) + exhausted = exhaust_huggingface_routing(unavailable, reason=reason) + committed = await self._persist_huggingface_routing_state( + exhausted, + expected_state=state, + ) + if not committed: + return None + return HuggingFaceRouteEvent( + status="exhausted", + previous_route=None, + route=response_route, + reason=reason, + ) + self._owned_providers.append(provider) + previous_state = state + resolved = resolve_automatic_huggingface_route(state, response_route) + resolved = observe_huggingface_cache( + resolved, + message.usage, + system=self._harness.config.system, + messages=self._harness.messages, + tools=self._harness.config.tools, + ) + previous_route = self._inference_provider + self._set_automatic_huggingface_route(response_route) + try: + committed = await self._persist_huggingface_routing_state( + resolved, + expected_state=previous_state, + ) + except Exception: + if self._is_current_huggingface_routing_state(previous_state): + self._set_automatic_huggingface_route(previous_route) + raise + if not committed: + return None + self._activate_runtime_provider(provider, provider_config) + return HuggingFaceRouteEvent( + status="changed", + previous_route=None, + route=response_route, + reason="automatic route resolved", + ) + + observed = observe_huggingface_cache( + state, + message.usage, + system=self._harness.config.system, + messages=self._harness.messages, + tools=self._harness.config.tools, + ) + if observed != state: + committed = await self._persist_huggingface_routing_state( + observed, + expected_state=state, + ) + if not committed: + return None + if state.phase != "exhausted" and observed.phase == "exhausted": + return HuggingFaceRouteEvent( + status="exhausted", + previous_route=None, + route=state.route, + reason=observed.last_reason or "eligible request budget exhausted", + ) + return None + + async def _reroute_huggingface_if_needed(self) -> HuggingFaceRouteEvent | None: + state = self._huggingface_routing_state + if ( + state is None + or self._inference_provider_mode != "automatic" + or state.phase != "reroute" + or state.route is None + ): + return None + + previous_route = state.route + try: + discovered = await discover_huggingface_routes(self.model) + except Exception as exc: # noqa: BLE001 - current route remains usable + if not self._is_current_huggingface_routing_state(state): + return None + reason = f"route discovery failed: {type(exc).__name__}: {exc}" + exhausted = exhaust_huggingface_routing(state, reason=reason) + committed = await self._persist_huggingface_routing_state( + exhausted, + expected_state=state, + ) + if not committed: + return None + return HuggingFaceRouteEvent( + status="exhausted", + previous_route=None, + route=previous_route, + reason=reason, + ) + + if not self._is_current_huggingface_routing_state(state): + return None + + failure_reason = state.last_reason or "warmed requests reported no cache reuse" + expected_state = state + for route in next_huggingface_routes(state, discovered): + try: + provider, provider_config = self._build_runtime_provider( + inference_provider=route, + ) + except ProviderConfigError: + state = mark_huggingface_route_unavailable(state, route) + continue + self._owned_providers.append(provider) + rerouted = reroute_huggingface_state(state, route, reason=failure_reason) + self._set_automatic_huggingface_route(route) + try: + committed = await self._persist_huggingface_routing_state( + rerouted, + expected_state=expected_state, + ) + except Exception: + if self._is_current_huggingface_routing_state(expected_state): + self._set_automatic_huggingface_route(previous_route) + raise + if not committed: + return None + self._activate_runtime_provider(provider, provider_config) + return HuggingFaceRouteEvent( + status="changed", + previous_route=previous_route, + route=route, + reason=failure_reason, + ) + + reason = "no untried live conversational routes were available" + exhausted = exhaust_huggingface_routing(state, reason=reason) + committed = await self._persist_huggingface_routing_state( + exhausted, + expected_state=expected_state, + ) + if not committed: + return None + return HuggingFaceRouteEvent( + status="exhausted", + previous_route=None, + route=previous_route, + reason=reason, + ) + + async def _persist_huggingface_routing_state( + self, + state: HuggingFaceRoutingState, + *, + expected_state: HuggingFaceRoutingState | None = None, + ) -> bool: + indexed_mode = self._indexed_inference_provider_mode() + await self.append_custom_entry(HF_CACHE_ROUTING_NAMESPACE, state.to_custom_data()) + if expected_state is not None and not self._is_current_huggingface_routing_state( + expected_state + ): + return False + self._huggingface_routing_state = state + if self._indexed_inference_provider_mode() != indexed_mode: + self._touch_session_index() + return True + + def _is_current_huggingface_routing_state(self, state: HuggingFaceRoutingState) -> bool: + return ( + self._inference_provider_mode == "automatic" + and self._huggingface_routing_state is state + ) + + def _set_automatic_huggingface_route(self, route: str | None) -> None: + self._inference_provider = route + self._inference_provider_mode = "automatic" + self._config = replace( + self._config, + inference_provider=route, + inference_provider_mode="automatic", + ) + + async def _handle_huggingface_agent_event( + self, + event: AgentEvent, + *, + context: AgentCallDiagnosticContext, + ) -> HuggingFaceRouteEvent | None: + try: + route_event: HuggingFaceRouteEvent | None = None + if isinstance(event, MessageEndEvent) and isinstance(event.message, AssistantMessage): + route_event = await self._observe_huggingface_response(event.message) + elif isinstance(event, AgentEndEvent): + route_event = await self._reroute_huggingface_if_needed() + except Exception as exc: # noqa: BLE001 - routing must not interrupt a completed turn + state = self._huggingface_routing_state + if ( + state is not None + and state.phase == "reroute" + and self._is_current_huggingface_routing_state(state) + ): + self._huggingface_routing_state = exhaust_huggingface_routing( + state, + reason=f"route evaluation failed: {type(exc).__name__}: {exc}", + ) + with suppress(Exception): + self._last_diagnostic_log_path = self._diagnostic_logger.log_exception( + context=context, + phase="huggingface_route_evaluation", + exc=exc, + ) + return None + if route_event is not None: + await self._extension_runtime.emit_event(route_event) + return route_event + def _invalidate_runtime_model_limits(self) -> None: self._runtime_model_limits = None self._runtime_model_limits_key = None @@ -1714,6 +2064,7 @@ async def resume(self, session_id: str) -> str: command_registry=self._config.command_registry, provider_name=provider_name, inference_provider=record.inference_provider, + inference_provider_mode=record.inference_provider_mode, provider_settings=self._provider_settings, runtime_provider_config=runtime_provider_config, auto_compact_token_threshold=self._auto_compact_token_threshold, @@ -1779,12 +2130,20 @@ async def new_session(self) -> str: ) inference_provider = _configured_inference_provider(runtime_provider_config, model) + inference_provider_mode: HuggingFaceRouteMode | None = ( + "explicit" + if provider_name == "huggingface" and inference_provider is not None + else "automatic" + if provider_name == "huggingface" + else None + ) record = ( manager.prepare_session( cwd=self.cwd, model=model, provider_name=provider_name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, ) if inference_provider is not None else manager.prepare_session( @@ -1803,6 +2162,7 @@ async def new_session(self) -> str: session_id=record.id, provider_name=provider_name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, provider_settings=self._provider_settings, runtime_provider_config=runtime_provider_config, thinking_level=thinking_level, @@ -1865,6 +2225,8 @@ async def _adopt_replacement( self._command_registry = replacement._command_registry self._provider_name = replacement._provider_name self._inference_provider = replacement._inference_provider + self._inference_provider_mode = replacement._inference_provider_mode + self._huggingface_routing_state = replacement._huggingface_routing_state self._provider_settings = replacement._provider_settings self._runtime_provider_config = replacement._runtime_provider_config self._resource_paths = replacement._resource_paths @@ -1950,6 +2312,7 @@ def ensure_session_indexed(self) -> None: model=self.model, provider_name=self.provider_name, inference_provider=self._inference_provider, + inference_provider_mode=self._indexed_inference_provider_mode(), session_id=self._config.session_id, ) self._config = replace(self._config, index_on_first_persist=False) @@ -2060,6 +2423,12 @@ async def prompt( ) await self._flush_pending_message_writes(context=context) + route_event = await self._handle_huggingface_agent_event( + AgentEndEvent(), + context=context, + ) + if route_event is not None: + yield route_event await self._refresh_runtime_model_limits() await self._try_auto_compact(context=context, phase="auto_compact_before_prompt") # id() values can be reused once earlier message objects are freed. @@ -2104,10 +2473,13 @@ async def prompt( ) if is_context_overflow_error(event.message): overflow_message = event.message + route_event = await self._handle_huggingface_agent_event(event, context=context) if isinstance(event, AgentEndEvent): yield SessionAgentEndEvent(messages=event.messages, will_retry=False) else: yield event + if route_event is not None: + yield route_event # Let frontends render the confirmed, expanded prompt before # session naming performs its separate provider request. if auto_name_message is not None: @@ -2152,6 +2524,10 @@ async def prompt( message=retry_event.message, ) ) + route_event = await self._handle_huggingface_agent_event( + retry_event, + context=context, + ) if isinstance(retry_event, AgentEndEvent): yield SessionAgentEndEvent( messages=retry_event.messages, @@ -2159,6 +2535,8 @@ async def prompt( ) else: yield retry_event + if route_event is not None: + yield route_event session_event_4 = AutoRetryEndEvent(success=True, attempt=1, final_error=None) await self._extension_runtime.emit_event(session_event_4) yield session_event_4 @@ -2184,6 +2562,12 @@ async def continue_(self) -> AsyncIterator[CodingSessionEvent]: """Continue the agent from restored state and persist new messages.""" context = self._diagnostic_context() await self._flush_pending_message_writes(context=context) + route_event = await self._handle_huggingface_agent_event( + AgentEndEvent(), + context=context, + ) + if route_event is not None: + yield route_event await self._refresh_runtime_model_limits() # id() values can be reused once earlier message objects are freed. self._ended_message_ids.clear() @@ -2205,10 +2589,13 @@ async def continue_(self) -> AsyncIterator[CodingSessionEvent]: phase="agent_loop", message=event.message, ) + route_event = await self._handle_huggingface_agent_event(event, context=context) if isinstance(event, AgentEndEvent): yield SessionAgentEndEvent(messages=event.messages, will_retry=False) else: yield event + if route_event is not None: + yield route_event await self._try_auto_compact(context=context, phase="auto_compact_after_continue") session_event_5 = AgentSettledEvent() await self._extension_runtime.emit_event(session_event_5) @@ -2419,14 +2806,7 @@ def _invalidate_context_usage_cache(self) -> None: async def _refresh_persisted_state(self, *, leaf_id: str | None) -> None: entries = await self._read_session_entries() self._state = SessionState.from_entries(entries, leaf_id=leaf_id) - if self._config.session_id is not None and self._config.session_manager is not None: - self._config.session_manager.touch_session( - self._config.session_id, - model=self.model, - provider_name=self.provider_name, - inference_provider=self._inference_provider, - preserve_inference_provider=False, - ) + self._touch_session_index() async def _read_session_entries(self) -> list[SessionEntry]: """Read stored entries, detaching roots imported from external history.""" @@ -2469,6 +2849,7 @@ def _index_current_session(self) -> None: model=self.model, provider_name=self.provider_name, inference_provider=self._inference_provider, + inference_provider_mode=self._indexed_inference_provider_mode(), session_id=self._config.session_id, ) @@ -3074,22 +3455,13 @@ def _create_runtime_provider( model: str, thinking_level: ThinkingLevel | None, inference_provider: str | None, - response_headers_observer: Callable[[Mapping[str, str]], None] | None = None, ) -> ClosableModelProvider: - if inference_provider is None and response_headers_observer is None: - return create_model_provider( - provider, - credential_store=credential_store, - model=model, - thinking_level=thinking_level, - ) if inference_provider is None: return create_model_provider( provider, credential_store=credential_store, model=model, thinking_level=thinking_level, - response_headers_observer=response_headers_observer, ) return create_model_provider( provider, @@ -3097,7 +3469,6 @@ def _create_runtime_provider( model=model, thinking_level=thinking_level, inference_provider=inference_provider, - response_headers_observer=response_headers_observer, ) diff --git a/src/tau_coding/session_manager.py b/src/tau_coding/session_manager.py index d5a9c6f7b..d1b71483f 100644 --- a/src/tau_coding/session_manager.py +++ b/src/tau_coding/session_manager.py @@ -7,6 +7,7 @@ from dataclasses import dataclass from pathlib import Path from time import time +from typing import Literal from uuid import uuid4 from pydantic import BaseModel, ConfigDict @@ -22,6 +23,8 @@ ) _SESSION_ID_PATTERN = re.compile(r"^[A-Za-z0-9](?:[A-Za-z0-9._-]*[A-Za-z0-9])?$") +type HuggingFaceRouteMode = Literal["automatic", "explicit"] + def validate_session_id(session_id: str) -> None: """Reject custom session ids that are unsafe as file names.""" @@ -50,6 +53,7 @@ class SessionRecordModel(BaseModel): model: str provider_name: str | None = None inference_provider: str | None = None + inference_provider_mode: HuggingFaceRouteMode | None = None title: str | None = None created_at: float updated_at: float @@ -68,6 +72,7 @@ class CodingSessionRecord: updated_at: float provider_name: str | None = None inference_provider: str | None = None + inference_provider_mode: HuggingFaceRouteMode | None = None @classmethod def from_model(cls, model: SessionRecordModel) -> CodingSessionRecord: @@ -82,6 +87,7 @@ def from_model(cls, model: SessionRecordModel) -> CodingSessionRecord: updated_at=model.updated_at, provider_name=model.provider_name, inference_provider=model.inference_provider, + inference_provider_mode=model.inference_provider_mode, ) def to_model(self) -> SessionRecordModel: @@ -96,6 +102,7 @@ def to_model(self) -> SessionRecordModel: updated_at=self.updated_at, provider_name=self.provider_name, inference_provider=self.inference_provider, + inference_provider_mode=self.inference_provider_mode, ) @@ -143,6 +150,7 @@ def create_session( model: str, provider_name: str | None = None, inference_provider: str | None = None, + inference_provider_mode: HuggingFaceRouteMode | None = None, title: str | None = None, session_id: str | None = None, ) -> CodingSessionRecord: @@ -152,6 +160,7 @@ def create_session( model=model, provider_name=provider_name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, title=title, session_id=session_id, ) @@ -165,6 +174,7 @@ def create_session_exclusive( model: str, provider_name: str | None = None, inference_provider: str | None = None, + inference_provider_mode: HuggingFaceRouteMode | None = None, title: str | None = None, session_id: str | None = None, ) -> CodingSessionRecord: @@ -174,6 +184,7 @@ def create_session_exclusive( model=model, provider_name=provider_name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, title=title, session_id=session_id, ) @@ -204,6 +215,7 @@ def prepare_session( model: str, provider_name: str | None = None, inference_provider: str | None = None, + inference_provider_mode: HuggingFaceRouteMode | None = None, title: str | None = None, session_id: str | None = None, ) -> CodingSessionRecord: @@ -225,6 +237,7 @@ def prepare_session( model=model, provider_name=provider_name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, title=title, created_at=now, updated_at=now, @@ -268,6 +281,7 @@ def touch_session( model: str | None = None, provider_name: str | None = None, inference_provider: str | None = None, + inference_provider_mode: HuggingFaceRouteMode | None = None, preserve_inference_provider: bool = True, title: str | None = None, ) -> CodingSessionRecord | None: @@ -284,6 +298,11 @@ def touch_session( inference_provider=( existing.inference_provider if preserve_inference_provider else inference_provider ), + inference_provider_mode=( + existing.inference_provider_mode + if preserve_inference_provider + else inference_provider_mode + ), title=title if title is not None else existing.title, created_at=existing.created_at, updated_at=time(), diff --git a/src/tau_coding/tui/adapter.py b/src/tau_coding/tui/adapter.py index f11835812..2bb0b204a 100644 --- a/src/tau_coding/tui/adapter.py +++ b/src/tau_coding/tui/adapter.py @@ -17,6 +17,7 @@ CodingSessionEvent, CompactionEndEvent, CompactionStartEvent, + HuggingFaceRouteEvent, QueueUpdateEvent, SessionAgentEndEvent, ) @@ -137,6 +138,9 @@ def apply(self, event: CodingSessionEvent) -> None: event.is_error, ) return + if isinstance(event, HuggingFaceRouteEvent): + self.state.add_item("status", event.display_text) + return if isinstance(event, CompactionStartEvent) and event.reason == "overflow": self.state.add_item("status", "… Context limit reached; compacting and retrying") return diff --git a/src/tau_coding/tui/app.py b/src/tau_coding/tui/app.py index 281e38f12..50cce64b0 100644 --- a/src/tau_coding/tui/app.py +++ b/src/tau_coding/tui/app.py @@ -7016,6 +7016,13 @@ async def run_tui_app( startup_message: str | None = None startup_error_notice: str | None = None runtime_provider_config: ProviderConfig | None = selection.provider + inference_provider_mode = ( + record.inference_provider_mode + if record is not None + and record.model == selection.model + and selection.provider.name == "huggingface" + else None + ) inference_provider = _startup_inference_provider(selection, record) try: provider = create_model_provider( @@ -7063,6 +7070,7 @@ async def run_tui_app( session_manager=manager, provider_name=selection.provider.name, inference_provider=inference_provider, + inference_provider_mode=inference_provider_mode, provider_settings=provider_settings, runtime_provider_config=runtime_provider_config, auto_compact_token_threshold=auto_compact_token_threshold, diff --git a/tests/test_coding_session.py b/tests/test_coding_session.py index bf93f1d3a..6e0dcca8f 100644 --- a/tests/test_coding_session.py +++ b/tests/test_coding_session.py @@ -3579,76 +3579,6 @@ async def test_session_compacts_and_retries_once_after_context_overflow( assert len(overflow_errors) == 1 -@pytest.mark.anyio -async def test_huggingface_session_pins_successful_automatic_route( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path -) -> None: - created: list[tuple[str | None, object | None]] = [] - - def create_provider( - provider_config: object, - *, - credential_store: FileCredentialStore | None = None, - model: str | None = None, - thinking_level: str | None = None, - inference_provider: str | None = None, - response_headers_observer: object | None = None, - ) -> SwitchableFakeProvider: - del credential_store, thinking_level - created.append((inference_provider, response_headers_observer)) - return SwitchableFakeProvider(provider_config) - - monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) - provider_config = OpenAICompatibleProviderConfig( - name="huggingface", - models=("zai-org/GLM-5.2",), - default_model="zai-org/GLM-5.2", - ) - manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) - record = manager.create_session( - cwd=tmp_path, - model="zai-org/GLM-5.2", - provider_name="huggingface", - ) - session = await CodingSession.load( - CodingSessionConfig( - provider=FakeProvider([]), - model="zai-org/GLM-5.2", - system="You are Tau.", - storage=JsonlSessionStorage(record.path), - cwd=tmp_path, - session_id=record.id, - session_manager=manager, - provider_name="huggingface", - provider_settings=ProviderSettings(providers=(provider_config,)), - runtime_provider_config=provider_config, - ) - ) - - observer = created[0][1] - assert callable(observer) - original_touch_session = manager.touch_session - active_provider = session._harness.config.provider - - def fail_touch_session(*args: object, **kwargs: object) -> None: - del args, kwargs - raise PermissionError("session index is read-only") - - monkeypatch.setattr(manager, "touch_session", fail_touch_session) - with pytest.raises(PermissionError, match="session index is read-only"): - observer({"X-Inference-Provider": "deepinfra"}) - - assert session.inference_provider is None - assert session._harness.config.provider is active_provider - - monkeypatch.setattr(manager, "touch_session", original_touch_session) - observer({"X-Inference-Provider": "deepinfra"}) - - assert session.inference_provider == "deepinfra" - assert created[-1][0] == "deepinfra" - assert manager.get_session(record.id).inference_provider == "deepinfra" # type: ignore[union-attr] - - @pytest.mark.anyio async def test_huggingface_session_re_resolves_pin_on_model_switch( monkeypatch: pytest.MonkeyPatch, tmp_path: Path diff --git a/tests/test_huggingface_routing.py b/tests/test_huggingface_routing.py new file mode 100644 index 000000000..5a983faff --- /dev/null +++ b/tests/test_huggingface_routing.py @@ -0,0 +1,938 @@ +from __future__ import annotations + +from asyncio import Event, create_task +from collections.abc import AsyncIterator +from dataclasses import replace +from pathlib import Path + +import httpx +import pytest + +from pi_event_helpers import assistant_done, assistant_start +from tau_agent.events import MessageEndEvent +from tau_agent.messages import AgentMessage, AssistantMessage, Usage, UserMessage +from tau_agent.provider_events import AssistantErrorEvent +from tau_agent.session import CustomEntry, JsonlSessionStorage, LeafEntry, MessageEntry +from tau_agent.tools import AgentTool +from tau_ai import CancellationToken, FakeProvider +from tau_ai.events import AssistantMessageEvent +from tau_coding import ( + CodingSession, + CodingSessionConfig, + OpenAICompatibleProviderConfig, + ProviderSettings, + SessionManager, + TauPaths, +) +from tau_coding import session as coding_session_module +from tau_coding.events import AgentSettledEvent, HuggingFaceRouteEvent, SessionAgentEndEvent +from tau_coding.huggingface_routing import ( + HF_CACHE_MAX_CANDIDATES, + HF_CACHE_MAX_ELIGIBLE_REQUESTS, + HF_CACHE_MIN_PROMPT_TOKENS, + HuggingFaceRoutingState, + discover_huggingface_routes, + next_huggingface_routes, + observe_huggingface_cache, +) + + +class _RouteAwareProvider: + def __init__( + self, + route: str | None, + responses: dict[str | None, list[AssistantMessage]], + calls: list[str | None], + ) -> None: + self.route = route + self.responses = responses + self.calls = calls + + def stream_response( + self, + *, + model: str, + system: str, + messages: list[AgentMessage], + tools: list[AgentTool], + signal: CancellationToken | None = None, + session_id: str | None = None, + ) -> AsyncIterator[AssistantMessageEvent]: + del system, messages, tools, signal, session_id + self.calls.append(self.route) + response = self.responses[self.route].pop(0) + + async def events() -> AsyncIterator[AssistantMessageEvent]: + yield assistant_start(model=model) + if response.stop_reason == "error": + yield AssistantErrorEvent(reason="error", error=response) + else: + yield assistant_done(response) + + return events() + + async def aclose(self) -> None: + pass + + +def _usage(*, cache_read: int = 0, reported: bool = False) -> Usage: + return Usage( + input=HF_CACHE_MIN_PROMPT_TOKENS - cache_read, + cache_read=cache_read, + cache_read_reported=reported, + total_tokens=HF_CACHE_MIN_PROMPT_TOKENS, + ) + + +def _messages(*values: str) -> tuple[UserMessage, ...]: + return tuple(UserMessage(content=value, timestamp=index) for index, value in enumerate(values)) + + +def _observe( + state: HuggingFaceRoutingState, + usage: Usage, + *messages: str, +) -> HuggingFaceRoutingState: + return observe_huggingface_cache( + state, + usage, + system="You are Tau.", + messages=_messages(*messages), + tools=(), + ) + + +async def _reroute_session( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> tuple[CodingSession, dict[str | None, _RouteAwareProvider]]: + model = "org/model" + provider_config = OpenAICompatibleProviderConfig( + name="huggingface", + models=(model,), + default_model=model, + ) + created: dict[str | None, _RouteAwareProvider] = {} + + def create_provider(*_args: object, **kwargs: object) -> _RouteAwareProvider: + route = kwargs.get("inference_provider") + assert route is None or isinstance(route, str) + provider = _RouteAwareProvider(route, {}, []) + created[route] = provider + return provider + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(tmp_path / "session.jsonl"), + cwd=tmp_path, + provider_name="huggingface", + inference_provider="scaleway", + inference_provider_mode="automatic", + runtime_provider_config=provider_config, + ) + ) + assert session._huggingface_routing_state is not None + session._huggingface_routing_state = session._huggingface_routing_state.model_copy( + update={"phase": "reroute"} + ) + return session, created + + +def test_huggingface_cache_policy_warms_then_reroutes_on_absent_telemetry() -> None: + state = HuggingFaceRoutingState.automatic("org/model", route="scaleway") + + cold = _observe(state, _usage(), "one") + first_probe = _observe(cold, _usage(), "one", "two") + reroute = _observe(first_probe, _usage(), "one", "two", "three") + + assert cold.phase == "evaluating" + assert cold.absent_probes == 0 + assert first_probe.phase == "evaluating" + assert first_probe.absent_probes == 1 + assert reroute.phase == "reroute" + assert reroute.absent_probes == 2 + assert reroute.zero_probes == 0 + assert reroute.eligible_requests == 3 + + assert _observe(reroute, _usage(), "one", "two", "three", "four") == reroute + + +def test_huggingface_cache_policy_distinguishes_zero_and_retains_positive_reuse() -> None: + state = HuggingFaceRoutingState.automatic("org/model", route="deepinfra") + first_hit = _observe(state, _usage(cache_read=2_048, reported=True), "one") + + cold = _observe(state, _usage(reported=True), "one") + zero = _observe(cold, _usage(reported=True), "one", "two") + retained = _observe(zero, _usage(cache_read=2_048, reported=True), "one", "two", "three") + later_miss = _observe(retained, _usage(), "one", "two", "three", "four") + + assert first_hit.phase == "retained" + assert cold.zero_probes == 0 + assert zero.zero_probes == 1 + assert zero.absent_probes == 0 + assert retained.phase == "retained" + assert retained.last_reason == "positive cache reuse reported" + assert later_miss == retained + + +def test_huggingface_cache_policy_ignores_short_and_incomparable_requests() -> None: + state = HuggingFaceRoutingState.automatic("org/model", route="deepinfra") + short_usage = Usage( + input=HF_CACHE_MIN_PROMPT_TOKENS - 1, + cache_read_reported=False, + total_tokens=HF_CACHE_MIN_PROMPT_TOKENS - 1, + ) + + short = _observe(state, short_usage, "one") + cold = _observe(short, _usage(), "one") + incomparable = _observe(cold, _usage(), "replacement") + first_probe = _observe(incomparable, _usage(), "replacement", "continued") + + assert short.phase == "evaluating" + assert short.eligible_requests == 0 + assert short.last_reason == "prompt below 4096-token eligibility threshold" + assert incomparable.absent_probes == 0 + assert incomparable.last_reason == "request context was not append-only comparable" + assert first_probe.absent_probes == 1 + + +def test_huggingface_cache_policy_exhaustion_is_terminal_and_bounded() -> None: + attempted = tuple(f"route-{index}" for index in range(HF_CACHE_MAX_CANDIDATES)) + state = HuggingFaceRoutingState.automatic("org/model", route=attempted[-1]).model_copy( + update={ + "attempted_routes": attempted, + "eligible_requests": HF_CACHE_MAX_ELIGIBLE_REQUESTS - 2, + } + ) + + cold = _observe(state, _usage(), "one") + exhausted = _observe(cold, _usage(), "one", "two") + unchanged = _observe(exhausted, _usage(cache_read=1, reported=True), "one", "two", "three") + + assert exhausted.phase == "exhausted" + assert exhausted.eligible_requests == HF_CACHE_MAX_ELIGIBLE_REQUESTS + assert "candidate budget" in (exhausted.last_reason or "") + assert unchanged == exhausted + + +@pytest.mark.anyio +async def test_huggingface_session_surfaces_budget_exhaustion(tmp_path: Path) -> None: + model = "org/model" + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(tmp_path / "session.jsonl"), + cwd=tmp_path, + provider_name="huggingface", + inference_provider="deepinfra", + inference_provider_mode="automatic", + ) + ) + session._huggingface_routing_state = HuggingFaceRoutingState.automatic( + model, + route="deepinfra", + ).model_copy( + update={ + "attempted_routes": ("scaleway", "novita", "deepinfra"), + "eligible_requests": HF_CACHE_MAX_ELIGIBLE_REQUESTS - 2, + } + ) + initial_state = session._huggingface_routing_state + assert ( + await session._observe_huggingface_response( + AssistantMessage(stop_reason="aborted", usage=_usage()) + ) + is None + ) + assert session._huggingface_routing_state is initial_state + + first_context = [UserMessage(content="one")] + session._harness.replace_messages(first_context) + assert await session._observe_huggingface_response(AssistantMessage(usage=_usage())) is None + + session._harness.replace_messages( + [*first_context, AssistantMessage(content="reply"), UserMessage(content="two")] + ) + event = await session._observe_huggingface_response(AssistantMessage(usage=_usage())) + + assert event is not None + assert event.status == "exhausted" + assert event.route == "deepinfra" + assert "budget" in event.reason + terminal_state = session._huggingface_routing_state + assert terminal_state is not None + assert terminal_state.phase == "exhausted" + assert await session._observe_huggingface_response(AssistantMessage(usage=_usage())) is None + assert session._huggingface_routing_state == terminal_state + await session.aclose() + + +@pytest.mark.anyio +async def test_huggingface_session_surfaces_unpinnable_automatic_route(tmp_path: Path) -> None: + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider([]), + model="org/model", + system="You are Tau.", + storage=JsonlSessionStorage(tmp_path / "session.jsonl"), + cwd=tmp_path, + provider_name="huggingface", + inference_provider_mode="automatic", + ) + ) + + event = await session._observe_huggingface_response( + AssistantMessage(response_provider="scaleway", usage=_usage()) + ) + + assert event is not None + assert event.status == "exhausted" + assert event.route == "scaleway" + assert event.reason == "scaleway could not be pinned locally" + assert "evaluation stopped" in (session.huggingface_routing_status or "") + await session.aclose() + + +@pytest.mark.anyio +async def test_huggingface_route_observation_failure_keeps_completed_response( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = "org/model" + provider_config = OpenAICompatibleProviderConfig( + name="huggingface", + models=(model,), + default_model=model, + ) + responses = { + None: [ + AssistantMessage( + content="completed response", + response_provider="deepinfra", + usage=_usage(), + ) + ] + } + calls: list[str | None] = [] + active_provider = _RouteAwareProvider(None, responses, calls) + staged_provider = _RouteAwareProvider("deepinfra", {}, calls) + + def create_provider(*_args: object, **kwargs: object) -> _RouteAwareProvider: + return ( + staged_provider if kwargs.get("inference_provider") == "deepinfra" else active_provider + ) + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) + record = manager.create_session( + cwd=tmp_path, + model=model, + provider_name="huggingface", + title="Routing failure test", + ) + session = await CodingSession.load( + CodingSessionConfig( + provider=active_provider, + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(record.path), + cwd=tmp_path, + session_id=record.id, + session_manager=manager, + provider_name="huggingface", + inference_provider_mode="automatic", + provider_settings=ProviderSettings(providers=(provider_config,)), + runtime_provider_config=provider_config, + ) + ) + original_touch_session = manager.touch_session + + def fail_route_touch_session(*args: object, **kwargs: object) -> object: + if kwargs.get("inference_provider") == "deepinfra": + raise PermissionError("session index is read-only") + return original_touch_session(*args, **kwargs) + + monkeypatch.setattr(manager, "touch_session", fail_route_touch_session) + events = await _collect(session.prompt("hello")) + + assert calls == [None] + assert any( + isinstance(event, SessionAgentEndEvent) + and any( + isinstance(message, AssistantMessage) and message.text == "completed response" + for message in event.messages + ) + for event in events + ) + assert isinstance(events[-1], AgentSettledEvent) + assert not any(isinstance(event, HuggingFaceRouteEvent) for event in events) + assert session.inference_provider is None + assert session._harness.config.provider is active_provider + assert session._huggingface_routing_state is not None + assert session._huggingface_routing_state.route is None + assert session._last_diagnostic_log_path is not None + await session.aclose() + + +@pytest.mark.anyio +async def test_huggingface_pending_reroute_failure_does_not_block_next_turn( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider( + [ + [ + assistant_start(model="org/model"), + assistant_done(AssistantMessage(content="completed response")), + ] + ] + ), + model="org/model", + system="You are Tau.", + storage=JsonlSessionStorage(tmp_path / "session.jsonl"), + cwd=tmp_path, + provider_name="huggingface", + inference_provider="scaleway", + inference_provider_mode="automatic", + ) + ) + assert session._huggingface_routing_state is not None + session._huggingface_routing_state = session._huggingface_routing_state.model_copy( + update={"phase": "reroute"} + ) + + async def fail_reroute() -> HuggingFaceRouteEvent | None: + raise PermissionError("session index is read-only") + + monkeypatch.setattr(session, "_reroute_huggingface_if_needed", fail_reroute) + + events = await _collect(session.prompt("hello")) + + assert any( + isinstance(event, MessageEndEvent) + and isinstance(event.message, AssistantMessage) + and event.message.text == "completed response" + for event in events + ) + assert isinstance(events[-1], AgentSettledEvent) + assert "evaluation stopped" in (session.huggingface_routing_status or "") + assert session._last_diagnostic_log_path is not None + await session.aclose() + + +@pytest.mark.anyio +async def test_huggingface_route_observation_precedes_consumer_teardown( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = "org/model" + responses = { + None: [ + AssistantMessage( + content="completed response", + response_provider="deepinfra", + usage=_usage(), + ) + ] + } + calls: list[str | None] = [] + + def create_provider(*_args: object, **kwargs: object) -> _RouteAwareProvider: + route = kwargs.get("inference_provider") + assert route is None or isinstance(route, str) + return _RouteAwareProvider(route, responses, calls) + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + provider_config = OpenAICompatibleProviderConfig( + name="huggingface", + models=(model,), + default_model=model, + ) + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(tmp_path / "session.jsonl"), + cwd=tmp_path, + provider_name="huggingface", + inference_provider_mode="automatic", + runtime_provider_config=provider_config, + ) + ) + + events = session.prompt("hello") + async for event in events: + if isinstance(event, MessageEndEvent) and isinstance(event.message, AssistantMessage): + break + await events.aclose() + + assert calls == [None] + assert session.inference_provider == "deepinfra" + assert session._huggingface_routing_state is not None + assert session._huggingface_routing_state.route == "deepinfra" + await session.aclose() + + +def test_huggingface_routing_state_round_trips_custom_entry_data() -> None: + state = _observe( + HuggingFaceRoutingState.automatic("org/model", route="deepinfra"), + _usage(reported=True), + "one", + ) + + restored = HuggingFaceRoutingState.from_custom_data(state.to_custom_data()) + + assert restored == state + assert HuggingFaceRoutingState.from_custom_data({"version": 999}) is None + + +@pytest.mark.anyio +async def test_huggingface_routing_state_survives_branch_then_resume(tmp_path: Path) -> None: + model = "org/model" + storage = JsonlSessionStorage(tmp_path / "session.jsonl") + user_entry = MessageEntry(message=UserMessage(content="one")) + assistant_entry = MessageEntry( + parent_id=user_entry.id, + message=AssistantMessage(content="reply"), + ) + await storage.append(user_entry) + await storage.append(assistant_entry) + await storage.append(LeafEntry(parent_id=assistant_entry.id, entry_id=assistant_entry.id)) + + manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) + record = manager.create_session( + cwd=tmp_path, + model=model, + provider_name="huggingface", + inference_provider="deepinfra", + inference_provider_mode="automatic", + title="Branch routing test", + ) + config = CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=storage, + cwd=tmp_path, + session_id=record.id, + session_manager=manager, + provider_name="huggingface", + inference_provider="deepinfra", + inference_provider_mode="automatic", + ) + session = await CodingSession.load(config) + retained = HuggingFaceRoutingState.automatic(model, route="deepinfra").model_copy( + update={ + "phase": "retained", + "eligible_requests": 3, + "last_reason": "positive cache reuse reported", + } + ) + await session._persist_huggingface_routing_state(retained) + + await session.branch_to_entry(assistant_entry.id) + await session.aclose() + + resumed = await CodingSession.load(replace(config, provider=FakeProvider([]))) + assert resumed._huggingface_routing_state == retained + await resumed.aclose() + + +@pytest.mark.anyio +async def test_huggingface_automatic_reset_survives_immediate_resume( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = "org/model" + provider_config = OpenAICompatibleProviderConfig( + name="huggingface", + models=(model,), + default_model=model, + ) + + def create_provider(*_args: object, **kwargs: object) -> _RouteAwareProvider: + route = kwargs.get("inference_provider") + assert route is None or isinstance(route, str) + return _RouteAwareProvider(route, {}, []) + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) + record = manager.create_session( + cwd=tmp_path, + model=model, + provider_name="huggingface", + ) + config = CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(record.path), + cwd=tmp_path, + session_id=record.id, + session_manager=manager, + provider_name="huggingface", + runtime_provider_config=provider_config, + ) + session = await CodingSession.load(config) + exhausted = HuggingFaceRoutingState.automatic(model).model_copy( + update={"phase": "exhausted", "last_reason": "previous automatic cycle exhausted"} + ) + await session._persist_huggingface_routing_state(exhausted) + exhausted_record = manager.get_session(record.id) + assert exhausted_record is not None + assert exhausted_record.inference_provider_mode == "automatic" + + session.set_inference_provider(None) + await session.aclose() + + reset_record = manager.get_session(record.id) + assert reset_record is not None + assert reset_record.inference_provider_mode is None + resumed = await CodingSession.load( + replace( + config, + provider=FakeProvider([]), + inference_provider=reset_record.inference_provider, + inference_provider_mode=reset_record.inference_provider_mode, + ) + ) + assert resumed.huggingface_routing_status == "automatic; waiting for route resolution" + await resumed.aclose() + + +def test_next_huggingface_routes_is_deterministic_and_budgeted() -> None: + state = HuggingFaceRoutingState.automatic("org/model", route="scaleway").model_copy( + update={"attempted_routes": ("scaleway", "baseten")} + ) + + candidates = next_huggingface_routes( + state, + ("novita", "deepinfra", "scaleway", "fireworks-ai", "baseten"), + ) + + assert candidates == ("deepinfra",) + + +@pytest.mark.anyio +async def test_discover_huggingface_routes_filters_live_conversational_mappings() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.params.get("expand[]") == "inferenceProviderMapping" + return httpx.Response( + 200, + json={ + "inferenceProviderMapping": { + "novita": {"status": "live", "task": "conversational"}, + "deepinfra": {"status": "live", "task": "conversational"}, + "offline": {"status": "staging", "task": "conversational"}, + "embeddings": {"status": "live", "task": "feature-extraction"}, + "malformed": "live", + } + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + routes = await discover_huggingface_routes( + "org/model", + client=client, + ) + + assert routes == ("deepinfra", "novita") + + +@pytest.mark.anyio +async def test_huggingface_session_reroutes_persists_and_resumes( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + responses = { + None: [ + AssistantMessage( + content="automatic response", + response_provider="scaleway", + usage=_usage(), + ) + ], + "scaleway": [ + AssistantMessage(stop_reason="error", error_message="temporary route error"), + AssistantMessage( + content="route mismatch", + response_provider="deepinfra", + usage=_usage(), + ), + AssistantMessage(content="first warmed miss", usage=_usage()), + AssistantMessage(content="second warmed miss", usage=_usage()), + ], + "deepinfra": [ + AssistantMessage(content="replacement warm-up", usage=_usage()), + AssistantMessage( + content="replacement cache hit", + usage=_usage(cache_read=2_048, reported=True), + ), + AssistantMessage(content="retained after resume", usage=_usage()), + ], + } + provider_calls: list[str | None] = [] + created_routes: list[str | None] = [] + + def create_provider(_provider_config: object, **kwargs: object) -> _RouteAwareProvider: + route = kwargs.get("inference_provider") + assert route is None or isinstance(route, str) + created_routes.append(route) + if route == "broken": + raise RuntimeError("route unavailable") + return _RouteAwareProvider(route, responses, provider_calls) + + discovery_calls: list[str] = [] + + async def discover_routes(model: str, **_kwargs: object) -> tuple[str, ...]: + discovery_calls.append(model) + return ("scaleway", "broken", "deepinfra") + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + monkeypatch.setattr(coding_session_module, "discover_huggingface_routes", discover_routes) + model = "org/model" + provider_config = OpenAICompatibleProviderConfig( + name="huggingface", + models=(model,), + default_model=model, + ) + manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) + record = manager.create_session( + cwd=tmp_path, + model=model, + provider_name="huggingface", + title="Cache routing test", + ) + config = CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(record.path), + cwd=tmp_path, + session_id=record.id, + session_manager=manager, + provider_name="huggingface", + inference_provider_mode="automatic", + provider_settings=ProviderSettings(providers=(provider_config,)), + runtime_provider_config=provider_config, + ) + session = await CodingSession.load(config) + + events = [] + for prompt in ("one", "two", "three", "four", "five", "six", "seven"): + events.extend([event async for event in session.prompt(prompt)]) + + route_events = [event for event in events if isinstance(event, HuggingFaceRouteEvent)] + assert [(event.previous_route, event.route) for event in route_events] == [ + (None, "scaleway"), + ("scaleway", "deepinfra"), + ] + assert provider_calls == [ + None, + "scaleway", + "scaleway", + "scaleway", + "scaleway", + "deepinfra", + "deepinfra", + ] + assert discovery_calls == [model] + assert "broken" in created_routes + current = manager.get_session(record.id) + assert current is not None + assert current.inference_provider == "deepinfra" + assert current.inference_provider_mode == "automatic" + snapshots = [ + HuggingFaceRoutingState.from_custom_data(entry.data) + for entry in session.state.custom_entries + if entry.namespace == "tau.huggingface-cache-routing" + ] + latest_snapshot = snapshots[-1] + assert latest_snapshot is not None + assert latest_snapshot.phase == "retained" + assert latest_snapshot.unavailable_routes == ("broken",) + status = session.handle_command("/session").message + assert status is not None + assert "Hugging Face cache routing: automatic; retained" in status + await session.aclose() + + resumed = await CodingSession.load( + replace( + config, + provider=FakeProvider([]), + inference_provider=current.inference_provider, + inference_provider_mode=current.inference_provider_mode, + ) + ) + assert "retained" in (resumed.huggingface_routing_status or "") + before_discovery = len(discovery_calls) + + await _collect(resumed.prompt("eight")) + + assert len(discovery_calls) == before_discovery + assert resumed.inference_provider == "deepinfra" + resumed.set_inference_provider(None) + await resumed.aclose() + + reset_record = manager.get_session(record.id) + assert reset_record is not None + restarted = await CodingSession.load( + replace( + config, + provider=FakeProvider([]), + inference_provider=reset_record.inference_provider, + inference_provider_mode=reset_record.inference_provider_mode, + ) + ) + assert restarted.inference_provider is None + assert restarted.huggingface_routing_status == "automatic; waiting for route resolution" + await restarted.aclose() + + +@pytest.mark.anyio +async def test_huggingface_explicit_route_stays_locked_until_reset( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + responses: dict[str | None, list[AssistantMessage]] = { + "scaleway": [AssistantMessage(content="explicit response", usage=_usage())], + None: [], + } + provider_calls: list[str | None] = [] + + def create_provider(_provider_config: object, **kwargs: object) -> _RouteAwareProvider: + route = kwargs.get("inference_provider") + assert route is None or isinstance(route, str) + return _RouteAwareProvider(route, responses, provider_calls) + + async def fail_discovery(*_args: object, **_kwargs: object) -> tuple[str, ...]: + raise AssertionError("explicit routes must not discover candidates") + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + monkeypatch.setattr(coding_session_module, "discover_huggingface_routes", fail_discovery) + model = "org/model" + provider_config = OpenAICompatibleProviderConfig( + name="huggingface", + models=(model,), + default_model=model, + ) + manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) + record = manager.create_session( + cwd=tmp_path, + model=model, + provider_name="huggingface", + inference_provider="scaleway", + title="Explicit route test", + ) + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider([]), + model=model, + system="You are Tau.", + storage=JsonlSessionStorage(record.path), + cwd=tmp_path, + session_id=record.id, + session_manager=manager, + provider_name="huggingface", + inference_provider="scaleway", + provider_settings=ProviderSettings(providers=(provider_config,)), + runtime_provider_config=provider_config, + ) + ) + + events = await _collect(session.prompt("one")) + + assert not any(isinstance(event, HuggingFaceRouteEvent) for event in events) + assert provider_calls == ["scaleway"] + assert session.huggingface_routing_status == "explicit route lock" + assert not any( + isinstance(entry, CustomEntry) and entry.namespace == "tau.huggingface-cache-routing" + for entry in session.state.custom_entries + ) + + assert session.set_inference_provider(None).startswith("automatic") + current = manager.get_session(record.id) + assert current is not None + assert current.inference_provider is None + assert current.inference_provider_mode is None + assert session.huggingface_routing_status == "automatic; waiting for route resolution" + + assert session.set_inference_provider("scaleway") == "scaleway" + current = manager.get_session(record.id) + assert current is not None + assert current.inference_provider_mode == "explicit" + assert session.huggingface_routing_status == "explicit route lock" + await session.aclose() + + +@pytest.mark.anyio +async def test_huggingface_explicit_route_wins_during_candidate_discovery( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + discovery_started = Event() + finish_discovery = Event() + + async def discover_routes(*_args: object, **_kwargs: object) -> tuple[str, ...]: + discovery_started.set() + await finish_discovery.wait() + return ("candidate",) + + monkeypatch.setattr(coding_session_module, "discover_huggingface_routes", discover_routes) + session, created = await _reroute_session(monkeypatch, tmp_path) + + reroute = create_task(session._reroute_huggingface_if_needed()) + await discovery_started.wait() + session.set_inference_provider("manual") + finish_discovery.set() + + assert await reroute is None + assert session.inference_provider == "manual" + assert session.huggingface_routing_status == "explicit route lock" + assert session._harness.config.provider is created["manual"] + await session.aclose() + + +@pytest.mark.anyio +async def test_huggingface_explicit_route_wins_during_automatic_state_write( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + async def discover_routes(*_args: object, **_kwargs: object) -> tuple[str, ...]: + return ("candidate",) + + monkeypatch.setattr(coding_session_module, "discover_huggingface_routes", discover_routes) + session, created = await _reroute_session(monkeypatch, tmp_path) + original_append = session.append_custom_entry + write_started = Event() + finish_write = Event() + + async def blocked_append(namespace: str, data: dict[str, object]) -> None: + write_started.set() + await finish_write.wait() + await original_append(namespace, data) # type: ignore[arg-type] + + monkeypatch.setattr(session, "append_custom_entry", blocked_append) + reroute = create_task(session._reroute_huggingface_if_needed()) + await write_started.wait() + session.set_inference_provider("manual") + finish_write.set() + + assert await reroute is None + assert session.inference_provider == "manual" + assert session.huggingface_routing_status == "explicit route lock" + assert session._harness.config.provider is created["manual"] + await session.aclose() + + +async def _collect(events: AsyncIterator[object]) -> list[object]: + return [event async for event in events] diff --git a/tests/test_rendering.py b/tests/test_rendering.py index 5a1b57e62..8d442a874 100644 --- a/tests/test_rendering.py +++ b/tests/test_rendering.py @@ -15,7 +15,7 @@ ToolExecutionUpdateEvent, ) from tau_agent.provider_events import TextDeltaEvent, ThinkingDeltaEvent -from tau_coding.events import AutoRetryStartEvent, QueueUpdateEvent +from tau_coding.events import AutoRetryStartEvent, HuggingFaceRouteEvent, QueueUpdateEvent from tau_coding.rendering import FinalTextRenderer, JsonEventRenderer, TranscriptRenderer @@ -168,6 +168,37 @@ def test_final_text_renderer_prints_errors_on_finish(capsys: pytest.CaptureFixtu assert "Error: provider failed" in capsys.readouterr().err +def test_renderers_surface_huggingface_route_changes( + capsys: pytest.CaptureFixture[str], +) -> None: + event = HuggingFaceRouteEvent( + status="changed", + previous_route="scaleway", + route="deepinfra", + reason="warmed requests reported no cache reuse", + ) + + TranscriptRenderer().render(event) + transcript = capsys.readouterr() + assert transcript.out == "" + assert "Hugging Face route changed: scaleway -> deepinfra" in transcript.err + assert "warmed requests reported no" in transcript.err + + FinalTextRenderer().render(event) + final_text = capsys.readouterr() + assert final_text.out == "" + assert "Hugging Face route changed: scaleway -> deepinfra" in final_text.err + + JsonEventRenderer().render(event) + assert json.loads(capsys.readouterr().out) == { + "type": "huggingface_route", + "status": "changed", + "previousRoute": "scaleway", + "route": "deepinfra", + "reason": "warmed requests reported no cache reuse", + } + + def test_json_renderer_emits_canonical_jsonl(capsys: pytest.CaptureFixture[str]) -> None: renderer = JsonEventRenderer() partial = AssistantMessage(content=[TextContent(text="hidden reasoning")]) diff --git a/tests/test_session_manager.py b/tests/test_session_manager.py index 8688a4211..65646b70b 100644 --- a/tests/test_session_manager.py +++ b/tests/test_session_manager.py @@ -19,11 +19,13 @@ def test_session_manager_creates_and_lists_sessions(tmp_path: Path) -> None: model="fake", provider_name="huggingface", inference_provider="deepinfra", + inference_provider_mode="automatic", title="Test session", ) assert record.provider_name == "huggingface" assert record.inference_provider == "deepinfra" + assert record.inference_provider_mode == "automatic" assert record.path.parent.parent == tmp_path / ".tau" / "sessions" assert "project-" in record.path.parent.name assert len(record.path.parent.name.rsplit("-", maxsplit=1)[-1]) == 6 @@ -263,6 +265,9 @@ def test_session_manager_touch_updates_metadata(tmp_path: Path) -> None: record.id, model="new-model", provider_name="new-provider", + inference_provider="deepinfra", + inference_provider_mode="explicit", + preserve_inference_provider=False, title="Updated", ) @@ -270,6 +275,8 @@ def test_session_manager_touch_updates_metadata(tmp_path: Path) -> None: assert updated.id == record.id assert updated.model == "new-model" assert updated.provider_name == "new-provider" + assert updated.inference_provider == "deepinfra" + assert updated.inference_provider_mode == "explicit" assert updated.title == "Updated" assert updated.updated_at >= record.updated_at assert manager.get_session(record.id) == updated diff --git a/tests/test_tau_ai.py b/tests/test_tau_ai.py index 226366326..063cd9def 100644 --- a/tests/test_tau_ai.py +++ b/tests/test_tau_ai.py @@ -2976,6 +2976,7 @@ def handler(request: httpx.Request) -> httpx.Response: assert usage.input == 20 # 30 prompt - 10 cached assert usage.output == 5 assert usage.cache_read == 10 + assert usage.cache_read_reported is True assert usage.cache_write == 0 assert usage.reasoning == 2 assert usage.total_tokens == 35 @@ -3084,6 +3085,7 @@ def handler(request: httpx.Request) -> httpx.Response: assert usage.input == 30 # 50 input - 12 cached - 8 written assert usage.output == 8 assert usage.cache_read == 12 + assert usage.cache_read_reported is True assert usage.cache_write == 8 assert usage.reasoning == 3 assert usage.total_tokens == 58 @@ -3216,6 +3218,7 @@ def handler(_request: httpx.Request) -> httpx.Response: assert usage.input == 15 assert usage.output == 4 assert usage.cache_read == 0 + assert usage.cache_read_reported is False assert usage.total_tokens == 19 @@ -3302,6 +3305,7 @@ def handler(_request: httpx.Request) -> httpx.Response: usage = events[-1].message.usage assert usage is not None assert usage.cache_read == 0 + assert usage.cache_read_reported is True assert usage.input == 40 diff --git a/tests/test_tui_adapter.py b/tests/test_tui_adapter.py index 61529a451..b0c8386a5 100644 --- a/tests/test_tui_adapter.py +++ b/tests/test_tui_adapter.py @@ -25,6 +25,7 @@ AutoRetryStartEvent, CompactionEndEvent, CompactionStartEvent, + HuggingFaceRouteEvent, QueueUpdateEvent, SessionAgentEndEvent, ) @@ -727,6 +728,28 @@ def test_tui_adapter_records_retry_and_queue_status() -> None: assert state.queued_follow_up == ("after",) +def test_tui_adapter_records_huggingface_route_status() -> None: + state = TuiState() + adapter = TuiEventAdapter(state) + + adapter.apply( + HuggingFaceRouteEvent( + status="exhausted", + previous_route="scaleway", + route="deepinfra", + reason="no untried live conversational routes were available", + ) + ) + + assert [(item.role, item.text) for item in state.items] == [ + ( + "status", + "Hugging Face route evaluation stopped on deepinfra: " + "no untried live conversational routes were available", + ) + ] + + def test_tui_adapter_records_assistant_error_and_aborted_message() -> None: state = TuiState(running=True, assistant_buffer="partial") adapter = TuiEventAdapter(state) diff --git a/website/content/guides/providers-and-models.md b/website/content/guides/providers-and-models.md index 73e9c4ae2..1f65d2745 100644 --- a/website/content/guides/providers-and-models.md +++ b/website/content/guides/providers-and-models.md @@ -151,10 +151,19 @@ and the inference provider selected by Hugging Face can vary over time and by account. For a new session without an explicit preference, Hugging Face initially routes -the model automatically. After the first successful response, Tau reads Hugging -Face's `x-inference-provider` response header and pins that backing provider for -the rest of the session. To choose the initial provider instead, add a per-model -`inference_providers` preference to `~/.tau/providers.json`: +the model automatically. After the first successful response, Tau pins the +backing provider reported by Hugging Face. Tau then checks cache reuse only for +append-only requests with at least 4,096 prompt tokens. The first eligible +request warms the route. If two later comparable requests report either no cache +telemetry or zero reused tokens, Tau tries another live conversational provider. +Any positive cache read retains the current route. A replacement starts cold and +may process or bill the full prompt, so evaluation stops after three routes or +nine eligible requests. Absent or zero cache telemetry means Tau did not observe +effective reported reuse; it does not prove that the provider has no internal +cache. + +To choose the provider instead, add a per-model `inference_providers` preference +to `~/.tau/providers.json`: ```json { @@ -172,8 +181,10 @@ the rest of the session. To choose the initial provider instead, add a per-model Use the exact provider suffix advertised for that model by Hugging Face. Tau sends `zai-org/GLM-5.2:deepinfra` on the wire and continues to display and store -the logical `zai-org/GLM-5.2` model. The pin survives resume; changing the -preference does not rewrite existing sessions. `/session` shows the active pin. +the logical `zai-org/GLM-5.2` model. The pin and automatic evaluation state +survive resume; changing the preference does not rewrite existing sessions. +Every automatic route change is shown in the transcript, and `/session` shows +the active route and evaluation status. Route selection is available through the external [`alejandro-ao/tau-huggingface`](https://github.com/alejandro-ao/tau-huggingface) extension rather than a built-in command. It requires Tau 0.3.10 or newer. Clone @@ -185,16 +196,17 @@ tau -e ./tau-huggingface ``` Then use `/route ` to select a route or `/route automatic` to reset -it. Switching models uses that model's configured pin or starts automatic -resolution again. +it. An explicit configured or extension-selected route stays locked and is not +evaluated automatically. Switching models uses that model's configured pin or +starts automatic resolution again. Transient failures retry on the same wire model, and stream failures are not retried after model output has started. Pinning can reduce cold prefix-cache misses caused by cross-provider routing, but cannot prevent eviction, TTL expiry, -or load balancing among workers within the chosen provider. Tau does not yet -fall back automatically from an unavailable pinned route: doing so also requires -a user-visible reroute event and durable reroute telemetry. Use the Hugging Face -extension or start a new automatic session to resolve another route. See +or load balancing among workers within the chosen provider. Short, failed, +aborted, route-mismatched, compacted, and non-comparable requests do not count as +cache-failure evidence. If no usable candidate remains, Tau keeps the current +route and stops evaluation. See [Configuration]({{< relref "../reference/configuration.md#provider-preferences" >}}). ### Moonshot AI API vs. Kimi Code diff --git a/website/content/guides/tui.md b/website/content/guides/tui.md index a5ae143ce..ec0e8286e 100644 --- a/website/content/guides/tui.md +++ b/website/content/guides/tui.md @@ -38,6 +38,10 @@ another prompt without starting a new session; empty failed provider turns are retained for diagnostics but are not replayed to the model as invalid conversation history. +For the built-in Hugging Face provider, an automatic route change or stopped +cache evaluation appears as a status row. `/session` reports the active backing +provider and its cache-routing state. + ## Cancelling and steering a run While the agent is working you don't have to wait: diff --git a/website/content/reference/configuration.md b/website/content/reference/configuration.md index 3c3265315..fc81bdacd 100644 --- a/website/content/reference/configuration.md +++ b/website/content/reference/configuration.md @@ -261,11 +261,17 @@ Provider preferences live in `~/.tau/providers.json`: Tau snapshots the selected suffix into new session metadata, retains it on resume, and sends only the suffixed wire model; ordinary model identity and catalog metadata remain unsuffixed. Without a preference, Tau starts with - automatic routing and pins the `x-inference-provider` reported by the first - successful response. `/session` reports the route; changing the active session - route is available in Tau 0.3.10+ through the external + automatic routing and pins the inference provider reported by the first + successful response. For append-only prompts of at least 4,096 tokens, it + treats the first request as a warm-up and changes routes after two later + responses report no positive cache reuse. Automatic evaluation is limited to + three routes and nine eligible requests, persists across resume, and stops on + the first positive cache read. `/session` reports the route and evaluation + status. An `inference_providers` preference remains explicitly locked; + changing the active session route is available in Tau 0.3.10+ through the external [`tau-huggingface`](https://github.com/alejandro-ao/tau-huggingface) extension; - clone it and launch Tau with `tau -e ./tau-huggingface`. + clone it and launch Tau with `tau -e ./tau-huggingface`. `/route automatic` + unlocks and resets automatic evaluation. `timeout_seconds` defaults to `60` (> 0); `max_retries` defaults to `2`; `max_retry_delay_seconds` defaults to `1` (both ≥ 0). Retries cover transient HTTP statuses (`408`, `409`, `425`, `429`, `5xx`),