diff --git a/src/tau_agent/loop.py b/src/tau_agent/loop.py
index e7681fcb3..60e03c723 100644
--- a/src/tau_agent/loop.py
+++ b/src/tau_agent/loop.py
@@ -231,10 +231,30 @@ async def _execute_tool(
)
if result.tool_call_id != tool_call.id:
- return result.model_copy(update={"tool_call_id": tool_call.id})
+ return _replace_tool_call_id(result, tool_call.id)
return result
+def _replace_tool_call_id(result: AgentToolResult, tool_call_id: str) -> AgentToolResult:
+ """Return a copy of a tool result with a provider-issued tool call id.
+
+ Tau runs on constrained Python environments that may expose Pydantic v1-style
+ models or local shims. Prefer Pydantic v2's ``model_copy`` when present, but
+ fall back to constructing a new result from the public fields.
+ """
+ model_copy = getattr(result, "model_copy", None)
+ if callable(model_copy):
+ return model_copy(update={"tool_call_id": tool_call_id})
+ return AgentToolResult(
+ tool_call_id=tool_call_id,
+ name=result.name,
+ ok=result.ok,
+ content=result.content,
+ error=result.error,
+ data=result.data,
+ )
+
+
def _unknown_tool_result(tool_call: ToolCall) -> AgentToolResult:
message = f"Unknown tool: {tool_call.name}"
return AgentToolResult(
diff --git a/src/tau_ai/__init__.py b/src/tau_ai/__init__.py
index fa327f915..c43ad8b47 100644
--- a/src/tau_ai/__init__.py
+++ b/src/tau_ai/__init__.py
@@ -23,6 +23,7 @@
ProviderToolCallEvent,
)
from tau_ai.fake import FakeProvider
+from tau_ai.models import ModelInfo, list_openai_compatible_models
from tau_ai.openai_codex import (
DEFAULT_OPENAI_CODEX_BASE_URL,
OpenAICodexConfig,
@@ -42,6 +43,7 @@
"DEFAULT_OPENAI_COMPATIBLE_TIMEOUT_SECONDS",
"DEFAULT_OPENAI_CODEX_BASE_URL",
"FakeProvider",
+ "ModelInfo",
"ModelProvider",
"OpenAICodexConfig",
"OpenAICodexCredentials",
@@ -56,5 +58,6 @@
"ProviderThinkingDeltaEvent",
"ProviderTextDeltaEvent",
"ProviderToolCallEvent",
+ "list_openai_compatible_models",
"openai_compatible_config_from_env",
]
diff --git a/src/tau_ai/anthropic.py b/src/tau_ai/anthropic.py
index edb0f0e13..2735774f9 100644
--- a/src/tau_ai/anthropic.py
+++ b/src/tau_ai/anthropic.py
@@ -3,6 +3,7 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Mapping
+from hashlib import sha1
from json import loads
from typing import Any
@@ -67,12 +68,16 @@ async def iterator() -> AsyncIterator[ProviderEvent]:
tools=tools,
thinking_budget_tokens=self._config.thinking_budget_tokens,
)
+ _sanitize_anthropic_payload_tool_ids(payload)
headers = {
**(dict(self._config.headers or {})),
"anthropic-version": ANTHROPIC_VERSION,
"content-type": "application/json",
- "x-api-key": self._config.api_key,
}
+ if self._config.auth_header.lower() == "authorization":
+ headers["Authorization"] = f"Bearer {self._config.api_key}"
+ else:
+ headers["x-api-key"] = self._config.api_key
url = f"{self._config.base_url.rstrip('/')}/messages"
attempt = 0
@@ -293,7 +298,7 @@ def _anthropic_message(message: AgentMessage) -> dict[str, JSONValue]:
content.append(
{
"type": "tool_use",
- "id": tool_call.id,
+ "id": _anthropic_tool_id(tool_call.id),
"name": tool_call.name,
"input": tool_call.arguments,
}
@@ -305,7 +310,7 @@ def _anthropic_message(message: AgentMessage) -> dict[str, JSONValue]:
"content": [
{
"type": "tool_result",
- "tool_use_id": message.tool_call_id,
+ "tool_use_id": _anthropic_tool_id(message.tool_call_id),
"content": message.content,
"is_error": not message.ok,
}
@@ -314,6 +319,45 @@ def _anthropic_message(message: AgentMessage) -> dict[str, JSONValue]:
raise TypeError(f"Unsupported message type: {type(message).__name__}")
+def _anthropic_tool_id(value: str) -> str:
+ """Return a tool-use id accepted by Anthropic's Messages API."""
+ cleaned = "".join(ch if ch.isalnum() or ch in {"_", "-"} else "_" for ch in value)
+ cleaned = cleaned.strip("_") or "tool_call"
+ if len(cleaned) <= 64:
+ return cleaned
+ suffix = "_" + sha1(value.encode("utf-8")).hexdigest()[:10]
+ return (cleaned[: 64 - len(suffix)].rstrip("_") or "tool_call") + suffix
+
+
+
+def _sanitize_anthropic_payload_tool_ids(payload: dict[str, JSONValue]) -> None:
+ """Defensively sanitize tool IDs in the final Anthropic payload."""
+ messages = payload.get("messages")
+ if not isinstance(messages, list):
+ return
+ id_map: dict[str, str] = {}
+ for message in messages:
+ if not isinstance(message, dict):
+ continue
+ content = message.get("content")
+ if not isinstance(content, list):
+ continue
+ for block in content:
+ if not isinstance(block, dict):
+ continue
+ if block.get("type") == "tool_use":
+ raw = block.get("id")
+ if isinstance(raw, str):
+ clean = _anthropic_tool_id(raw)
+ id_map[raw] = clean
+ block["id"] = clean
+ elif block.get("type") == "tool_result":
+ raw = block.get("tool_use_id")
+ if isinstance(raw, str):
+ block["tool_use_id"] = id_map.get(raw, _anthropic_tool_id(raw))
+
+
+
def _anthropic_tool(tool: AgentTool) -> dict[str, JSONValue]:
return {
"name": tool.name,
diff --git a/src/tau_ai/env.py b/src/tau_ai/env.py
index 0bba65fe0..b54f897d1 100644
--- a/src/tau_ai/env.py
+++ b/src/tau_ai/env.py
@@ -5,6 +5,7 @@
from collections.abc import Mapping
from dataclasses import dataclass
from os import environ
+from typing import Literal
DEFAULT_OPENAI_COMPATIBLE_BASE_URL = "https://api.openai.com/v1"
DEFAULT_ANTHROPIC_BASE_URL = "https://api.anthropic.com/v1"
@@ -12,6 +13,8 @@
DEFAULT_OPENAI_COMPATIBLE_MAX_RETRIES = 2
DEFAULT_OPENAI_COMPATIBLE_MAX_RETRY_DELAY_SECONDS = 1.0
+AnthropicThinkingType = Literal["adaptive", "disabled"]
+
@dataclass(frozen=True, slots=True)
class OpenAICompatibleConfig:
@@ -38,6 +41,8 @@ class AnthropicConfig:
max_retries: int = DEFAULT_OPENAI_COMPATIBLE_MAX_RETRIES
max_retry_delay_seconds: float = DEFAULT_OPENAI_COMPATIBLE_MAX_RETRY_DELAY_SECONDS
thinking_budget_tokens: int | None = None
+ thinking_type: AnthropicThinkingType | None = None
+ auth_header: str = "x-api-key"
def openai_compatible_config_from_env(
diff --git a/src/tau_ai/models.py b/src/tau_ai/models.py
new file mode 100644
index 000000000..a8068ed41
--- /dev/null
+++ b/src/tau_ai/models.py
@@ -0,0 +1,93 @@
+"""Model discovery for OpenAI-compatible providers.
+
+Some OpenAI-compatible endpoints (for example Nebius Token Factory) expose a
+``GET /v1/models`` listing that can be expanded with a ``verbose`` query
+parameter. Tau uses this to populate a provider's model list dynamically at
+build time instead of hardcoding a catalog.
+"""
+
+from __future__ import annotations
+
+from collections.abc import Mapping
+from dataclasses import dataclass
+from json import loads
+from typing import Any
+
+import httpx
+
+from tau_ai.env import OpenAICompatibleConfig
+
+
+@dataclass(frozen=True, slots=True)
+class ModelInfo:
+ """One model advertised by an OpenAI-compatible ``/models`` endpoint."""
+
+ id: str
+ context_window: int | None = None
+
+
+async def list_openai_compatible_models(
+ config: OpenAICompatibleConfig,
+ *,
+ verbose: bool = False,
+ client: httpx.AsyncClient | None = None,
+) -> tuple[ModelInfo, ...]:
+ """List models from an OpenAI-compatible ``/models`` endpoint.
+
+ When ``verbose`` is true, the ``verbose=true`` query parameter is sent so
+ providers that support it (such as Nebius Token Factory) return the full
+ model catalog with metadata. The response is parsed tolerantly: only the
+ ``id`` of each entry in ``data`` is required, and an optional integer
+ context-window field is extracted when present.
+ """
+ headers: dict[str, str] = {**(dict(config.headers or {}))}
+ headers.setdefault("Authorization", f"Bearer {config.api_key}")
+ params: dict[str, str] = {}
+ if verbose:
+ params["verbose"] = "true"
+ url = f"{config.base_url.rstrip('/')}/models"
+
+ owns_client = client is None
+ http_client = client or httpx.AsyncClient(timeout=config.timeout_seconds)
+ try:
+ response = await http_client.get(url, headers=headers, params=params)
+ response.raise_for_status()
+ payload = loads(response.content)
+ finally:
+ if owns_client:
+ await http_client.aclose()
+
+ data = _data_array(payload)
+ models: list[ModelInfo] = []
+ seen: set[str] = set()
+ for entry in data:
+ model_id = _model_id(entry)
+ if model_id is None or model_id in seen:
+ continue
+ seen.add(model_id)
+ models.append(ModelInfo(id=model_id, context_window=_context_window(entry)))
+ return tuple(models)
+
+
+def _data_array(payload: object) -> list[Mapping[str, Any]]:
+ if not isinstance(payload, Mapping):
+ return []
+ data = payload.get("data")
+ if not isinstance(data, list):
+ return []
+ return [item for item in data if isinstance(item, Mapping)]
+
+
+def _model_id(entry: Mapping[str, Any]) -> str | None:
+ model_id = entry.get("id")
+ if isinstance(model_id, str) and model_id.strip():
+ return model_id.strip()
+ return None
+
+
+def _context_window(entry: Mapping[str, Any]) -> int | None:
+ for field_name in ("context_window", "context_length", "max_context_length"):
+ value = entry.get(field_name)
+ if isinstance(value, int) and not isinstance(value, bool) and value > 0:
+ return value
+ return None
diff --git a/src/tau_ai/openai_codex.py b/src/tau_ai/openai_codex.py
index a01df3e80..2b6439f11 100644
--- a/src/tau_ai/openai_codex.py
+++ b/src/tau_ai/openai_codex.py
@@ -4,6 +4,7 @@
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
from dataclasses import dataclass
+from hashlib import sha1
from json import JSONDecodeError, dumps, loads
from platform import machine, release, system
from typing import Any
@@ -318,15 +319,16 @@ def _messages_to_responses_input(messages: list[AgentMessage]) -> list[JSONValue
call_id, item_id = _split_tool_call_id(tool_call.id)
item: dict[str, JSONValue] = {
"type": "function_call",
- "call_id": call_id,
- "name": tool_call.name,
+ "call_id": _codex_call_id(call_id),
+ "name": tool_call.name or "tool",
"arguments": dumps(tool_call.arguments),
}
if item_id:
- item["id"] = item_id
+ item["id"] = _codex_item_id(item_id)
items.append(item)
elif isinstance(message, ToolResultMessage):
call_id, _item_id = _split_tool_call_id(message.tool_call_id)
+ call_id = _codex_call_id(call_id)
items.append(
{
"type": "function_call_output",
@@ -337,6 +339,27 @@ def _messages_to_responses_input(messages: list[AgentMessage]) -> list[JSONValue
return items
+def _codex_identifier(value: str, *, fallback: str) -> str:
+ """Return a Codex Responses identifier safe for replayed transcript items."""
+ cleaned = "".join(ch if ch.isalnum() or ch in {"_", "-"} else "_" for ch in value)
+ cleaned = cleaned.strip("_") or fallback
+ if len(cleaned) <= 64:
+ return cleaned
+ suffix = "_" + sha1(value.encode("utf-8")).hexdigest()[:10]
+ return (cleaned[: 64 - len(suffix)].rstrip("_") or fallback) + suffix
+
+
+
+def _codex_call_id(value: str) -> str:
+ return _codex_identifier(value, fallback="call")
+
+
+
+def _codex_item_id(value: str) -> str:
+ return _codex_identifier(value, fallback="item")
+
+
+
def _tool_to_codex(tool: AgentTool) -> dict[str, JSONValue]:
return {
"type": "function",
diff --git a/src/tau_ai/openai_compatible.py b/src/tau_ai/openai_compatible.py
index 4b60a3634..e5414b341 100644
--- a/src/tau_ai/openai_compatible.py
+++ b/src/tau_ai/openai_compatible.py
@@ -11,6 +11,7 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Callable, Mapping
+from hashlib import sha1
from json import JSONDecodeError, dumps, loads
from typing import Any, Protocol
@@ -625,8 +626,8 @@ def _messages_to_responses_input(
items.append(
{
"type": "function_call",
- "call_id": tool_call.id,
- "name": tool_call.name,
+ "call_id": _responses_call_id(tool_call.id),
+ "name": tool_call.name or "tool",
"arguments": dumps(tool_call.arguments),
}
)
@@ -634,13 +635,29 @@ def _messages_to_responses_input(
items.append(
{
"type": "function_call_output",
- "call_id": message.tool_call_id,
+ "call_id": _responses_call_id(message.tool_call_id),
"output": message.content,
}
)
return items
+def _responses_call_id(value: str) -> str:
+ """Return a Responses API call_id accepted by OpenAI-compatible backends.
+
+ Provider transcripts can persist foreign tool-call IDs with separators or
+ long opaque suffixes. Normalize separators and add a short hash suffix when
+ truncating so function_call/function_call_output pairs remain stable.
+ """
+ cleaned = "".join(ch if ch.isalnum() or ch in {"_", "-"} else "_" for ch in value)
+ cleaned = cleaned.strip("_") or "call"
+ if len(cleaned) <= 64:
+ return cleaned
+ suffix = "_" + sha1(value.encode("utf-8")).hexdigest()[:10]
+ return (cleaned[: 64 - len(suffix)].rstrip("_") or "call") + suffix
+
+
+
def _tool_to_responses(tool: AgentTool) -> dict[str, JSONValue]:
return {
"type": "function",
@@ -769,7 +786,7 @@ def _message_to_openai(message: AgentMessage) -> dict[str, JSONValue]:
return {
"role": "tool",
"tool_call_id": message.tool_call_id,
- "name": message.name,
+ "name": message.name or "tool",
"content": message.content,
}
@@ -790,7 +807,7 @@ def _tool_call_to_openai(tool_call: ToolCall) -> dict[str, JSONValue]:
"id": tool_call.id,
"type": "function",
"function": {
- "name": tool_call.name,
+ "name": tool_call.name or "tool",
"arguments": dumps(tool_call.arguments),
},
}
diff --git a/src/tau_coding/branch_summary.py b/src/tau_coding/branch_summary.py
index adc1392ea..c345bb9af 100644
--- a/src/tau_coding/branch_summary.py
+++ b/src/tau_coding/branch_summary.py
@@ -158,8 +158,12 @@ def _format_assistant_summary_source(message: AssistantMessage) -> str:
def _format_tool_call_arguments(arguments: Mapping[str, object]) -> str:
+ public_arguments = {
+ key: value for key, value in arguments.items() if key != "_raw_arguments"
+ }
return ", ".join(
- f"{key}={json.dumps(value, sort_keys=True)}" for key, value in sorted(arguments.items())
+ f"{key}={json.dumps(value, sort_keys=True)}"
+ for key, value in sorted(public_arguments.items())
)
diff --git a/src/tau_coding/cli.py b/src/tau_coding/cli.py
index 5a76c9b7d..cd8576bca 100644
--- a/src/tau_coding/cli.py
+++ b/src/tau_coding/cli.py
@@ -432,11 +432,11 @@ def _provider_credential_status(
credential_reader: CredentialReader | None,
) -> str:
if provider.credential_name and credential_reader is not None:
- if provider_kind(provider) == "openai-codex":
+ if provider_kind(provider) == "openai-codex" or provider.name == "github-copilot":
get_oauth = getattr(credential_reader, "get_oauth", None)
if get_oauth is not None and get_oauth(provider.credential_name) is not None:
return f"stored:{provider.credential_name}"
- elif credential_reader.get(provider.credential_name):
+ if credential_reader.get(provider.credential_name):
return f"stored:{provider.credential_name}"
if environ.get(provider.api_key_env):
return f"env:{provider.api_key_env}"
diff --git a/src/tau_coding/context_window.py b/src/tau_coding/context_window.py
index 3d3d90f5c..37a2dc65e 100644
--- a/src/tau_coding/context_window.py
+++ b/src/tau_coding/context_window.py
@@ -238,7 +238,8 @@ def serialize_messages_for_compaction(messages: tuple[AgentMessage, ...]) -> str
if message.tool_calls:
lines.append("
{escaped}
' + + +@dataclass(frozen=True, slots=True) +class GitHubCopilotDeviceFlow: + """GitHub device-code login state for Copilot.""" + + device_code: str + user_code: str + verification_uri: str + interval: int + expires_in: int + + +async def login_github_copilot( + *, + on_auth: AuthCallback, + on_prompt: PromptCallback | None = None, + on_progress: ProgressCallback | None = None, + client: httpx.AsyncClient | None = None, +) -> OAuthCredential: + """Run GitHub Copilot device OAuth and return refreshable credentials. + + The stored refresh value is the GitHub OAuth token. The stored access value + is the short-lived Copilot API token, refreshed on demand via GitHub's + ``/copilot_internal/v2/token`` endpoint. + """ + enterprise_domain = "" + if on_prompt is not None: + raw_domain = await on_prompt( + OAuthPrompt( + message="GitHub Enterprise URL/domain (blank for github.com):", + placeholder="company.ghe.com", + ) + ) + enterprise_domain = normalize_github_domain(raw_domain) + domain = enterprise_domain or "github.com" + + owns_client = client is None + http_client = client or httpx.AsyncClient(timeout=30) + try: + device = await start_github_copilot_device_flow(domain, client=http_client) + on_auth( + OAuthAuthInfo( + url=device.verification_uri, + instructions=f"Enter code: {device.user_code}", + ) + ) + github_token = await poll_github_copilot_device_flow( + domain, + device, + client=http_client, + ) + on_progress and on_progress("Exchanging GitHub token for Copilot token...") + credential = await refresh_github_copilot_token( + github_token, + enterprise_domain=enterprise_domain, + client=http_client, + ) + on_progress and on_progress("Fetching Copilot model availability...") + return credential + finally: + if owns_client: + await http_client.aclose() + + +async def start_github_copilot_device_flow( + domain: str = "github.com", + *, + client: httpx.AsyncClient | None = None, +) -> GitHubCopilotDeviceFlow: + """Start GitHub's device-code OAuth flow for Copilot.""" + owns_client = client is None + http_client = client or httpx.AsyncClient(timeout=30) + try: + response = await http_client.post( + f"https://{domain}/login/device/code", + data={"client_id": GITHUB_COPILOT_CLIENT_ID, "scope": "read:user"}, + headers={ + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + "User-Agent": GITHUB_COPILOT_HEADERS["User-Agent"], + }, + ) + response.raise_for_status() + raw = response.json() + return GitHubCopilotDeviceFlow( + device_code=_required_string(raw, "device_code", action="device flow"), + user_code=_required_string(raw, "user_code", action="device flow"), + verification_uri=_required_string(raw, "verification_uri", action="device flow"), + interval=int(raw.get("interval") or 5), + expires_in=int(raw.get("expires_in") or 900), + ) + finally: + if owns_client: + await http_client.aclose() + + +async def poll_github_copilot_device_flow( + domain: str, + device: GitHubCopilotDeviceFlow, + *, + client: httpx.AsyncClient | None = None, +) -> str: + """Poll GitHub's device-code endpoint until the GitHub token is available.""" + owns_client = client is None + http_client = client or httpx.AsyncClient(timeout=30) + interval = max(device.interval, 5) + deadline = time.monotonic() + device.expires_in + try: + while time.monotonic() < deadline: + await asyncio.sleep(interval) + response = await http_client.post( + f"https://{domain}/login/oauth/access_token", + data={ + "client_id": GITHUB_COPILOT_CLIENT_ID, + "device_code": device.device_code, + "grant_type": "urn:ietf:params:oauth:grant-type:device_code", + }, + headers={ + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + "User-Agent": GITHUB_COPILOT_HEADERS["User-Agent"], + }, + ) + response.raise_for_status() + raw = response.json() + token = raw.get("access_token") + if isinstance(token, str) and token.strip(): + return token.strip() + error = raw.get("error") + if error in {"authorization_pending", None}: + continue + if error == "slow_down": + interval += 5 + continue + description = raw.get("error_description") + suffix = f": {description}" if isinstance(description, str) else "" + raise OAuthError(f"GitHub device flow failed: {error}{suffix}") + finally: + if owns_client: + await http_client.aclose() + raise OAuthError("GitHub device flow timed out") + + +async def refresh_github_copilot_token( + github_token: str, + *, + enterprise_domain: str = "", + client: httpx.AsyncClient | None = None, +) -> OAuthCredential: + """Exchange a GitHub OAuth token for a short-lived Copilot API token.""" + domain = enterprise_domain or "github.com" + owns_client = client is None + http_client = client or httpx.AsyncClient(timeout=30) + try: + response = await http_client.get( + f"https://api.{domain}/copilot_internal/v2/token", + headers={ + **GITHUB_COPILOT_HEADERS, + "Accept": "application/json", + "Authorization": f"Bearer {github_token}", + }, + ) + response.raise_for_status() + raw = response.json() + access = _required_string(raw, "token", action="copilot token") + expires_at = raw.get("expires_at") + if not isinstance(expires_at, int | float) or isinstance(expires_at, bool): + raise OAuthError("Missing Copilot token expiry") + return OAuthCredential( + access=access, + refresh=github_token, + expires=int(expires_at * 1000) - 5 * 60 * 1000, + account_id=enterprise_domain or "github.com", + ) + finally: + if owns_client: + await http_client.aclose() + + +def github_copilot_base_url(token: str, enterprise_domain: str = "") -> str: + """Return the Copilot API base URL for a token or enterprise domain.""" + marker = "proxy-ep=" + if marker in token: + host = token.split(marker, 1)[1].split(";", 1)[0] + if host: + return "https://" + host.replace("proxy.", "api.", 1) + if enterprise_domain and enterprise_domain != "github.com": + return f"https://copilot-api.{enterprise_domain}" + return "https://api.individual.githubcopilot.com" + + +def normalize_github_domain(value: str | None) -> str: + """Normalize a user-entered GitHub Enterprise URL/domain.""" + stripped = (value or "").strip() + if not stripped: + return "" + if "://" not in stripped: + stripped = "https://" + stripped + parsed = urlparse(stripped) + return parsed.hostname or "" + + +def _required_string(raw: dict[str, Any], field: str, *, action: str) -> str: + value = raw.get(field) + if not isinstance(value, str) or not value.strip(): + raise OAuthError(f"Missing {field} in GitHub Copilot {action} response") + return value.strip() diff --git a/src/tau_coding/provider_catalog.py b/src/tau_coding/provider_catalog.py index d4b61afdd..699d75d25 100644 --- a/src/tau_coding/provider_catalog.py +++ b/src/tau_coding/provider_catalog.py @@ -2,6 +2,7 @@ from __future__ import annotations +from collections.abc import Mapping from dataclasses import dataclass from typing import Literal @@ -10,6 +11,24 @@ ProviderKind = Literal["openai-compatible", "anthropic", "openai-codex"] +@dataclass(frozen=True, slots=True) +class ThinkingMode: + """A canonical Tau thinking level's provider-specific behavior.""" + + api_value: str | None = None + label: str | None = None + + +@dataclass(frozen=True, slots=True) +class ProviderModelOverride: + """Built-in model behavior that differs from its provider defaults.""" + + kind: ProviderKind | None = None + thinking_modes: Mapping[ThinkingLevel, ThinkingMode] | None = None + thinking_default: ThinkingLevel | None = None + always_thinking: bool = False + + @dataclass(frozen=True, slots=True) class ProviderCatalogEntry: """A built-in provider Tau can present during login.""" @@ -28,6 +47,8 @@ class ProviderCatalogEntry: thinking_models: tuple[str, ...] = () thinking_default: ThinkingLevel | None = None thinking_parameter: ThinkingParameter | None = None + model_overrides: dict[str, ProviderModelOverride] | None = None + dynamic_models: bool = False BUILTIN_PROVIDER_CATALOG: tuple[ProviderCatalogEntry, ...] = ( @@ -145,6 +166,63 @@ class ProviderCatalogEntry: thinking_default="medium", thinking_parameter="anthropic.thinking", ), + ProviderCatalogEntry( + name="github-copilot", + display_name="GitHub Copilot", + kind="openai-compatible", + base_url="https://api.individual.githubcopilot.com", + api_key_env="GITHUB_COPILOT_TOKEN", + credential_name="github-copilot", + models=( + "gpt-5.5", + "gpt-5.4", + "gpt-5.4-mini", + "gpt-5.3-codex", + "gpt-5-mini", + "claude-sonnet-5", + "claude-sonnet-4.6", + "claude-sonnet-4.5", + "claude-opus-4.8", + "claude-opus-4.7", + "claude-opus-4.6", + "claude-haiku-4.5", + "gemini-3.5-flash", + "gemini-3.1-pro-preview", + "gemini-3-flash-preview", + "gemini-2.5-pro", + ), + default_model="gpt-5.5", + docs_url="https://docs.github.com/copilot", + context_windows={ + "gpt-5.5": 272_000, + "gpt-5.4": 272_000, + "gpt-5.4-mini": 400_000, + "gpt-5.3-codex": 400_000, + "gpt-5-mini": 400_000, + "claude-sonnet-5": 1_000_000, + "claude-sonnet-4.6": 1_000_000, + "claude-sonnet-4.5": 1_000_000, + "claude-opus-4.8": 1_000_000, + "claude-opus-4.7": 1_000_000, + "claude-opus-4.6": 1_000_000, + "claude-haiku-4.5": 200_000, + "gemini-3.5-flash": 1_048_576, + "gemini-3.1-pro-preview": 1_048_576, + "gemini-3-flash-preview": 1_048_576, + "gemini-2.5-pro": 1_048_576, + }, + thinking_levels=("off", "low", "medium", "high", "xhigh"), + thinking_models=( + "gpt-5.5", + "gpt-5.4", + "gpt-5.4-mini", + "gpt-5.3-codex", + "gpt-5-mini", + ), + thinking_default="medium", + thinking_parameter="reasoning_effort", + dynamic_models=True, + ), ProviderCatalogEntry( name="openrouter", display_name="OpenRouter", @@ -235,6 +313,165 @@ class ProviderCatalogEntry: thinking_default="medium", thinking_parameter="reasoning_effort", ), + ProviderCatalogEntry( + name="deepseek", + display_name="DeepSeek", + kind="openai-compatible", + base_url="https://api.deepseek.com/v1", + api_key_env="DEEPSEEK_API_KEY", + credential_name="deepseek", + models=( + "deepseek-v4-flash", + "deepseek-v4-pro", + ), + default_model="deepseek-v4-flash", + docs_url="https://api-docs.deepseek.com", + context_windows={ + "deepseek-v4-flash": 1_048_576, + "deepseek-v4-pro": 1_048_576, + }, + thinking_levels=("off", "low", "medium", "high", "xhigh"), + thinking_models=( + "deepseek-v4-flash", + "deepseek-v4-pro", + ), + thinking_default="high", + thinking_parameter="reasoning_effort", + ), + ProviderCatalogEntry( + name="opencode-go", + display_name="OpenCode Go", + kind="openai-compatible", + base_url="https://opencode.ai/zen/go/v1", + api_key_env="OPENCODE_GO_API_KEY", + credential_name="opencode-go", + models=( + "glm-5.2", + "glm-5.1", + "kimi-k2.7-code", + "kimi-k2.6", + "deepseek-v4-pro", + "deepseek-v4-flash", + "mimo-v2.5", + "mimo-v2.5-pro", + "minimax-m3", + "minimax-m2.7", + "qwen3.7-max", + "qwen3.7-plus", + "qwen3.6-plus", + ), + default_model="deepseek-v4-pro", + docs_url="https://opencode.ai/docs/go", + context_windows={ + "glm-5.2": 1_000_000, + "glm-5.1": 202_752, + "kimi-k2.7-code": 262_144, + "kimi-k2.6": 262_144, + "deepseek-v4-pro": 1_000_000, + "deepseek-v4-flash": 1_000_000, + "mimo-v2.5": 1_000_000, + "mimo-v2.5-pro": 1_048_576, + "minimax-m3": 1_000_000, + "minimax-m2.7": 204_800, + "qwen3.7-max": 1_000_000, + "qwen3.7-plus": 1_000_000, + "qwen3.6-plus": 1_000_000, + }, + thinking_levels=("off", "low", "medium", "high", "xhigh"), + thinking_models=( + "glm-5.2", + "deepseek-v4-pro", + "deepseek-v4-flash", + "minimax-m3", + "qwen3.7-max", + "qwen3.7-plus", + "qwen3.6-plus", + ), + thinking_default="medium", + thinking_parameter="reasoning_effort", + model_overrides={ + "glm-5.2": ProviderModelOverride( + thinking_modes={ + "high": ThinkingMode(api_value="high"), + "xhigh": ThinkingMode(api_value="max", label="max"), + }, + thinking_default="high", + ), + "glm-5.1": ProviderModelOverride(always_thinking=True), + "kimi-k2.7-code": ProviderModelOverride(always_thinking=True), + "kimi-k2.6": ProviderModelOverride(always_thinking=True), + "deepseek-v4-pro": ProviderModelOverride( + thinking_modes={ + "high": ThinkingMode(api_value="high"), + "xhigh": ThinkingMode(api_value="max", label="max"), + }, + thinking_default="high", + ), + "deepseek-v4-flash": ProviderModelOverride( + thinking_modes={ + "high": ThinkingMode(api_value="high"), + "xhigh": ThinkingMode(api_value="max", label="max"), + }, + thinking_default="high", + ), + "mimo-v2.5": ProviderModelOverride(always_thinking=True), + "mimo-v2.5-pro": ProviderModelOverride(always_thinking=True), + "minimax-m3": ProviderModelOverride( + kind="anthropic", + thinking_modes={ + "off": ThinkingMode(api_value="disabled"), + "high": ThinkingMode(api_value="adaptive", label="on"), + }, + thinking_default="high", + ), + "minimax-m2.7": ProviderModelOverride(kind="anthropic", always_thinking=True), + "qwen3.7-max": ProviderModelOverride( + kind="anthropic", + thinking_modes={ + "off": ThinkingMode(api_value="disabled"), + "low": ThinkingMode(), + "medium": ThinkingMode(), + "high": ThinkingMode(), + "xhigh": ThinkingMode(), + }, + thinking_default="medium", + ), + "qwen3.7-plus": ProviderModelOverride( + kind="anthropic", + thinking_modes={ + "off": ThinkingMode(api_value="disabled"), + "low": ThinkingMode(), + "medium": ThinkingMode(), + "high": ThinkingMode(), + "xhigh": ThinkingMode(), + }, + thinking_default="medium", + ), + "qwen3.6-plus": ProviderModelOverride( + kind="anthropic", + thinking_modes={ + "off": ThinkingMode(api_value="disabled"), + "low": ThinkingMode(), + "medium": ThinkingMode(), + "high": ThinkingMode(), + "xhigh": ThinkingMode(), + }, + thinking_default="medium", + ), + }, + ), + ProviderCatalogEntry( + name="nebius", + display_name="Nebius Token Factory", + kind="openai-compatible", + base_url="https://api.tokenfactory.nebius.com/v1", + api_key_env="NEBIUS_TOKEN_FACTORY_API_KEY", + credential_name="nebius", + models=(), + default_model="", + docs_url="https://docs.tokenfactory.nebius.com", + dynamic_models=True, + ), ) @@ -244,3 +481,13 @@ def builtin_provider_entry(name: str) -> ProviderCatalogEntry | None: if entry.name == name: return entry return None + + +def catalog_model_override(provider_name: str, model: str | None) -> ProviderModelOverride | None: + """Return built-in metadata for a provider/model pair.""" + if model is None: + return None + entry = builtin_provider_entry(provider_name) + if entry is None or entry.model_overrides is None: + return None + return entry.model_overrides.get(model) diff --git a/src/tau_coding/provider_config.py b/src/tau_coding/provider_config.py index 526447a6c..8c2cfbeb8 100644 --- a/src/tau_coding/provider_config.py +++ b/src/tau_coding/provider_config.py @@ -11,6 +11,8 @@ from tempfile import NamedTemporaryFile from typing import Any, Protocol +import httpx + from tau_ai import ( DEFAULT_ANTHROPIC_BASE_URL, DEFAULT_OPENAI_CODEX_BASE_URL, @@ -19,13 +21,26 @@ DEFAULT_OPENAI_COMPATIBLE_TIMEOUT_SECONDS, AnthropicConfig, OpenAICompatibleConfig, + list_openai_compatible_models, ) -from tau_ai.env import DEFAULT_OPENAI_COMPATIBLE_BASE_URL +from tau_ai.env import DEFAULT_OPENAI_COMPATIBLE_BASE_URL, AnthropicThinkingType from tau_coding.credentials import FileCredentialStore, credentials_path +from tau_coding.oauth import ( + github_copilot_base_url, + oauth_credential_is_expired, + refresh_github_copilot_token, +) from tau_coding.paths import TauPaths -from tau_coding.provider_catalog import BUILTIN_PROVIDER_CATALOG, ProviderKind +from tau_coding.provider_catalog import ( + BUILTIN_PROVIDER_CATALOG, + ProviderKind, + ThinkingMode, + builtin_provider_entry, + catalog_model_override, +) from tau_coding.thinking import ( DEFAULT_THINKING_LEVEL, + THINKING_LEVELS, ThinkingLevel, ThinkingParameter, anthropic_thinking_budget_for_level, @@ -67,6 +82,7 @@ class OpenAICompatibleProviderConfig: thinking_models: tuple[str, ...] = () thinking_default: ThinkingLevel | None = None thinking_parameter: ThinkingParameter | None = None + dynamic_models: bool = False def __post_init__(self) -> None: _validate_provider_numbers( @@ -103,6 +119,7 @@ def to_json(self) -> dict[str, Any]: "thinking_models": list(self.thinking_models), "thinking_default": self.thinking_default, "thinking_parameter": self.thinking_parameter, + "dynamic_models": self.dynamic_models, } @@ -114,8 +131,8 @@ class AnthropicProviderConfig: base_url: str = DEFAULT_ANTHROPIC_BASE_URL api_key_env: str = "ANTHROPIC_API_KEY" credential_name: str | None = "anthropic" - models: tuple[str, ...] = ("claude-sonnet-4-6",) - default_model: str = "claude-sonnet-4-6" + models: tuple[str, ...] = ("claude-fable-5", "claude-sonnet-5", "claude-sonnet-4-6") + default_model: str = "claude-sonnet-5" context_windows: dict[str, int] = field(default_factory=dict) headers: dict[str, str] = field(default_factory=dict) timeout_seconds: float = DEFAULT_OPENAI_COMPATIBLE_TIMEOUT_SECONDS @@ -334,6 +351,7 @@ def provider_config_from_catalog_entry(name: str) -> ProviderConfig: thinking_models=entry.thinking_models, thinking_default=entry.thinking_default, thinking_parameter=entry.thinking_parameter, + dynamic_models=entry.dynamic_models, ) raise ProviderConfigError(f"Unknown built-in provider: {name}") @@ -491,6 +509,20 @@ def upsert_provider( return updated +def _replace_provider(settings: ProviderSettings, provider: ProviderConfig) -> ProviderSettings: + """Return settings with an exact provider replacement, without built-in model merging.""" + providers_by_name = {item.name: item for item in settings.providers} + providers_by_name[provider.name] = provider + providers = tuple(providers_by_name[name] for name in sorted(providers_by_name)) + updated = ProviderSettings( + default_provider=settings.default_provider, + providers=providers, + scoped_models=settings.scoped_models, + ) + updated.get_provider(settings.default_provider) + return updated + + def _with_builtin_catalog_models( settings: ProviderSettings, *, @@ -668,14 +700,60 @@ def provider_thinking_levels( model: str | None = None, ) -> tuple[ThinkingLevel, ...]: """Return thinking levels supported by a provider/model pair.""" + selected_model = model or provider.default_model + override = catalog_model_override(provider.name, selected_model) + if override is not None: + if override.always_thinking: + return () + if override.thinking_modes is not None: + return tuple(level for level in THINKING_LEVELS if level in override.thinking_modes) if provider.thinking_levels is None: return () - selected_model = model or provider.default_model if provider.thinking_models and selected_model not in provider.thinking_models: return () return provider.thinking_levels +def provider_thinking_is_always_on( + provider: ProviderConfig, + *, + model: str | None = None, +) -> bool: + """Return whether built-in metadata declares reasoning as always enabled.""" + selected_model = model or provider.default_model + override = catalog_model_override(provider.name, selected_model) + return override.always_thinking if override is not None else False + + +def provider_thinking_level_label( + provider: ProviderConfig, + level: str, + *, + model: str | None = None, +) -> str: + """Return the provider-facing display label for a canonical Tau level.""" + normalized = normalize_thinking_level(level) + mode = _thinking_mode(provider, model=model, level=normalized) + return mode.label if mode is not None and mode.label is not None else normalized + + +def provider_thinking_level_from_label( + provider: ProviderConfig, + value: str, + *, + model: str | None = None, +) -> ThinkingLevel: + """Resolve a provider-facing input label to a canonical Tau level.""" + selected_model = model or provider.default_model + override = catalog_model_override(provider.name, selected_model) + if override is not None and override.thinking_modes is not None: + label = value.strip().lower() + for level, mode in override.thinking_modes.items(): + if mode.label == label: + return level + return normalize_thinking_level(value) + + def provider_thinking_unavailable_reason( provider: ProviderConfig, *, @@ -683,6 +761,8 @@ def provider_thinking_unavailable_reason( ) -> str | None: """Explain why a provider/model pair has no configurable thinking modes.""" selected_model = model or provider.default_model + if provider_thinking_is_always_on(provider, model=selected_model): + return f"Reasoning is always enabled for {selected_model}" if provider.thinking_levels is None: if isinstance(provider, OpenAICodexProviderConfig): return ( @@ -705,6 +785,10 @@ def provider_default_thinking_level( levels = provider_thinking_levels(provider, model=model) if not levels: return None + selected_model = model or provider.default_model + override = catalog_model_override(provider.name, selected_model) + if override is not None and override.thinking_default in levels: + return override.thinking_default if provider.thinking_default in levels: return provider.thinking_default if DEFAULT_THINKING_LEVEL in levels: @@ -741,18 +825,137 @@ def openai_compatible_config_from_provider( ) +async def ensure_dynamic_provider_models( + settings: ProviderSettings, + *, + provider_name: str, + paths: TauPaths | None = None, + credential_store: FileCredentialStore | None = None, + client: httpx.AsyncClient | None = None, +) -> ProviderSettings: + """Populate a dynamic provider's model list at build time. + + Built-in providers flagged ``dynamic_models`` (such as Nebius Token Factory) + start with an empty model catalog. When Tau has usable credentials for the + selected provider, this fetches the live model list from the provider's + ``/models`` endpoint (with ``verbose=true``), persists it through the normal + provider-settings path, and returns the updated settings. It is best-effort: + any network, auth, or parse error leaves the settings unchanged so startup + never fails because of a model listing problem. + """ + entry = builtin_provider_entry(provider_name) + if entry is None or not entry.dynamic_models: + return settings + try: + provider = settings.get_provider(provider_name) + except ProviderConfigError: + return settings + if not isinstance(provider, OpenAICompatibleProviderConfig): + return settings + + store = credential_store or FileCredentialStore(credentials_path(paths) if paths else None) + if not provider_has_usable_credentials(provider, credential_reader=store): + return settings + + try: + if provider.name == "github-copilot": + provider = await _github_copilot_provider_for_model_listing( + provider, + credential_store=store, + ) + runtime_config = openai_compatible_config_from_provider( + provider, credential_reader=store + ) + models = await list_openai_compatible_models( + runtime_config, verbose=True, client=client + ) + except Exception: + return settings + + if not models: + return settings + + model_ids = tuple(model.id for model in models) + context_windows = { + **dict(provider.context_windows), + **{ + model.id: model.context_window + for model in models + if model.context_window is not None + }, + } + default_model = provider.default_model if provider.default_model in model_ids else model_ids[0] + updated_provider = replace( + provider, + models=model_ids, + default_model=default_model, + context_windows=context_windows, + ) + updated = _replace_provider(settings, updated_provider) + save_provider_settings(updated, paths) + return updated + + +async def _github_copilot_provider_for_model_listing( + provider: OpenAICompatibleProviderConfig, + *, + credential_store: FileCredentialStore, +) -> OpenAICompatibleProviderConfig: + """Return a Copilot provider config suitable for calling its OpenAI-compatible API.""" + headers = {**dict(provider.headers), **_github_copilot_headers()} + credential_name = provider.credential_name + if not credential_name: + return replace(provider, headers=headers) + credential = credential_store.get_oauth(credential_name) + if credential is None: + return replace(provider, headers=headers) + if oauth_credential_is_expired(credential): + credential = await refresh_github_copilot_token( + credential.refresh, + enterprise_domain=credential.account_id if credential.account_id != "github.com" else "", + ) + credential_store.set_oauth(credential_name, credential) + return replace( + provider, + base_url=github_copilot_base_url(credential.access, credential.account_id), + headers=headers, + ) + + +def _github_copilot_headers() -> dict[str, str]: + return { + "User-Agent": "GitHubCopilotChat/0.35.0", + "Editor-Version": "vscode/1.107.0", + "Editor-Plugin-Version": "copilot-chat/0.35.0", + "Copilot-Integration-Id": "vscode-chat", + "openai-intent": "conversation-panel", + } + + def anthropic_config_from_provider( provider: AnthropicProviderConfig, *, credential_reader: CredentialReader | None = None, + model: str | None = None, thinking_level: ThinkingLevel | None = None, ) -> AnthropicConfig: """Build Anthropic runtime config from durable settings.""" api_key = _api_key_from_provider(provider, credential_reader=credential_reader) - thinking_budget_tokens = _anthropic_thinking_budget_from_provider( + thinking_type = _anthropic_thinking_type_from_provider( provider, + model=model, thinking_level=thinking_level, ) + thinking_budget_tokens = ( + None + if thinking_type is not None + else _anthropic_thinking_budget_from_provider( + provider, + model=model, + thinking_level=thinking_level, + ) + ) + auth_header = "authorization" if provider.name == "github-copilot" else "x-api-key" return AnthropicConfig( api_key=api_key, base_url=provider.base_url.rstrip("/"), @@ -761,6 +964,8 @@ def anthropic_config_from_provider( max_retries=provider.max_retries, max_retry_delay_seconds=provider.max_retry_delay_seconds, thinking_budget_tokens=thinking_budget_tokens, + thinking_type=thinking_type, + auth_header=auth_header, ) @@ -780,11 +985,11 @@ def provider_has_usable_credentials( ) -> bool: """Return whether Tau can attempt calls for this provider without prompting setup.""" if provider.credential_name and credential_reader is not None: - if isinstance(provider, OpenAICodexProviderConfig): + if isinstance(provider, OpenAICodexProviderConfig) or provider.name == "github-copilot": get_oauth = getattr(credential_reader, "get_oauth", None) if get_oauth is not None and get_oauth(provider.credential_name) is not None: return True - elif credential_reader.get(provider.credential_name): + if credential_reader.get(provider.credential_name): return True return bool(environ.get(provider.api_key_env)) @@ -813,31 +1018,81 @@ def _reasoning_effort_from_provider( f"Thinking mode {normalized} is not available for " f"{provider.name}:{selected_model}. Available modes: {available}" ) - return reasoning_effort_for_level(normalized) + api_value = _thinking_api_value(provider, model=model, level=normalized) + return api_value if api_value is not None else reasoning_effort_for_level(normalized) def _anthropic_thinking_budget_from_provider( provider: AnthropicProviderConfig, *, + model: str | None = None, thinking_level: ThinkingLevel | None, ) -> int | None: if thinking_level is None or provider.thinking_parameter != "anthropic.thinking": return None - levels = provider_thinking_levels(provider) + levels = provider_thinking_levels(provider, model=model) if not levels: return None normalized = normalize_thinking_level(thinking_level) if normalized not in levels: + selected_model = model or provider.default_model available = ", ".join(levels) raise ProviderConfigError( f"Thinking mode {normalized} is not available for " - f"{provider.name}:{provider.default_model}. Available modes: {available}" + f"{provider.name}:{selected_model}. Available modes: {available}" ) return anthropic_thinking_budget_for_level(normalized) +def _anthropic_thinking_type_from_provider( + provider: AnthropicProviderConfig, + *, + model: str | None, + thinking_level: ThinkingLevel | None, +) -> AnthropicThinkingType | None: + if thinking_level is None or provider.thinking_parameter != "anthropic.thinking": + return None + normalized = normalize_thinking_level(thinking_level) + levels = provider_thinking_levels(provider, model=model) + if not levels: + return None + if normalized not in levels: + selected_model = model or provider.default_model + available = ", ".join(levels) + raise ProviderConfigError( + f"Thinking mode {normalized} is not available for " + f"{provider.name}:{selected_model}. Available modes: {available}" + ) + api_value = _thinking_api_value(provider, model=model, level=normalized) + if api_value == "adaptive": + return "adaptive" + if api_value == "disabled": + return "disabled" + return None + + +def _thinking_api_value( + provider: ProviderConfig, + *, + model: str | None, + level: ThinkingLevel, +) -> str | None: + mode = _thinking_mode(provider, model=model, level=level) + return mode.api_value if mode is not None else None + + +def _thinking_mode( + provider: ProviderConfig, *, model: str | None, level: ThinkingLevel +) -> ThinkingMode | None: + selected_model = model or provider.default_model + override = catalog_model_override(provider.name, selected_model) + if override is None or override.thinking_modes is None: + return None + return override.thinking_modes.get(level) + + def _provider_from_json(data: object) -> ProviderConfig: if not isinstance(data, dict): raise ProviderConfigError("Provider entries must be JSON objects") @@ -850,8 +1105,15 @@ def _provider_from_json(data: object) -> ProviderConfig: credential_name = _optional_string( data.get("credential_name"), f"providers[{name}].credential_name" ) - models = _string_tuple(data.get("models"), f"providers[{name}].models") - default_model = _string(data.get("default_model"), f"providers[{name}].default_model") + dynamic_models = bool(data.get("dynamic_models", False)) + if dynamic_models: + models = _optional_string_tuple(data.get("models"), f"providers[{name}].models") + default_model = _emptyable_string( + data.get("default_model"), f"providers[{name}].default_model" + ) + else: + models = _string_tuple(data.get("models"), f"providers[{name}].models") + default_model = _string(data.get("default_model"), f"providers[{name}].default_model") context_windows = _context_window_dict( data.get("context_windows", {}), f"providers[{name}].context_windows" ) @@ -883,7 +1145,7 @@ def _provider_from_json(data: object) -> ProviderConfig: thinking_parameter = _optional_thinking_parameter( data.get("thinking_parameter"), f"providers[{name}].thinking_parameter" ) - if default_model not in models: + if default_model and default_model not in models: models = (*models, default_model) if provider_type == "anthropic": return AnthropicProviderConfig( @@ -937,6 +1199,7 @@ def _provider_from_json(data: object) -> ProviderConfig: thinking_models=thinking_models, thinking_default=thinking_default, thinking_parameter=thinking_parameter, + dynamic_models=dynamic_models, ) @@ -949,6 +1212,11 @@ def _api_key_from_provider( credential = credential_reader.get(provider.credential_name) if credential: return credential + get_oauth = getattr(credential_reader, "get_oauth", None) + if get_oauth is not None: + oauth_credential = get_oauth(provider.credential_name) + if oauth_credential is not None: + return oauth_credential.access api_key = environ.get(provider.api_key_env) if api_key: @@ -1039,6 +1307,15 @@ def _optional_string(value: object, field_name: str) -> str | None: return value.strip() +def _emptyable_string(value: object, field_name: str) -> str: + """Parse a string field that may be empty (used by dynamic model catalogs).""" + if value is None: + return "" + if not isinstance(value, str): + raise ProviderConfigError(f"Provider field must be a string: {field_name}") + return value.strip() + + def _string(value: object, field_name: str) -> str: if not isinstance(value, str) or not value.strip(): raise ProviderConfigError(f"Provider field must be a non-empty string: {field_name}") diff --git a/src/tau_coding/provider_runtime.py b/src/tau_coding/provider_runtime.py index 89f04d3bb..528dbc9e1 100644 --- a/src/tau_coding/provider_runtime.py +++ b/src/tau_coding/provider_runtime.py @@ -2,23 +2,33 @@ from __future__ import annotations +import asyncio +from collections.abc import AsyncIterator +from dataclasses import replace from os import environ from typing import Protocol +from tau_agent.messages import AgentMessage +from tau_agent.tools import AgentTool from tau_ai import ( AnthropicProvider, + CancellationToken, ModelProvider, OpenAICodexConfig, OpenAICodexCredentials, OpenAICodexProvider, OpenAICompatibleProvider, + ProviderEvent, ) from tau_coding.credentials import FileCredentialStore, OAuthCredential from tau_coding.oauth import ( account_id_from_access_token, + github_copilot_base_url, oauth_credential_is_expired, + refresh_github_copilot_token, refresh_openai_codex_token, ) +from tau_coding.provider_catalog import catalog_model_override from tau_coding.provider_config import ( AnthropicProviderConfig, OpenAICodexProviderConfig, @@ -48,11 +58,26 @@ def create_model_provider( ) -> ClosableModelProvider: """Create a runtime model provider from durable provider settings.""" credentials = credential_store or FileCredentialStore() + selected_model = model or provider.default_model + if provider.name == "github-copilot": + return GitHubCopilotCredentialRefreshingProvider( + provider, + credential_store=credentials, + thinking_level=thinking_level, + ) + override = catalog_model_override(provider.name, selected_model) + if ( + provider.name != "github-copilot" + and override is not None + and override.kind == "anthropic" + ): + provider = _anthropic_provider_config_for_model(provider, selected_model) if isinstance(provider, AnthropicProviderConfig): return AnthropicProvider( anthropic_config_from_provider( provider, credential_reader=credentials, + model=selected_model, thinking_level=thinking_level, ) ) @@ -79,12 +104,173 @@ def create_model_provider( openai_compatible_config_from_provider( provider, credential_reader=credentials, - model=model, + model=selected_model, thinking_level=thinking_level, ) ) +def _github_copilot_provider_config( + provider: ProviderConfig, + *, + credential_store: FileCredentialStore, +) -> ProviderConfig: + """Refresh and adapt GitHub Copilot OAuth settings for runtime calls.""" + credential_name = provider.credential_name + if credential_name: + credential = credential_store.get_oauth(credential_name) + if credential is not None: + credential = _refresh_github_copilot_if_needed( + credential_name, + credential, + credential_store=credential_store, + ) + base_url = github_copilot_base_url(credential.access, credential.account_id) + return replace( + provider, + base_url=base_url, + headers={**dict(provider.headers), **_github_copilot_headers()}, + ) + return replace(provider, headers={**dict(provider.headers), **_github_copilot_headers()}) + + +def _refresh_github_copilot_if_needed( + credential_name: str, + credential: OAuthCredential, + *, + credential_store: FileCredentialStore, +) -> OAuthCredential: + """Refresh Copilot credentials when provider setup is outside an event loop.""" + if not oauth_credential_is_expired(credential): + return credential + try: + asyncio.get_running_loop() + except RuntimeError: + refreshed = asyncio.run( + refresh_github_copilot_token( + credential.refresh, + enterprise_domain=credential.account_id if credential.account_id != "github.com" else "", + ) + ) + credential_store.set_oauth(credential_name, refreshed) + return refreshed + return credential + + +def _github_copilot_headers() -> dict[str, str]: + return { + "User-Agent": "GitHubCopilotChat/0.35.0", + "Editor-Version": "vscode/1.107.0", + "Editor-Plugin-Version": "copilot-chat/0.35.0", + "Copilot-Integration-Id": "vscode-chat", + "openai-intent": "conversation-panel", + } + + +class GitHubCopilotCredentialRefreshingProvider: + """GitHub Copilot provider wrapper with request-time OAuth refresh.""" + + def __init__( + self, + provider: ProviderConfig, + *, + credential_store: FileCredentialStore, + thinking_level: ThinkingLevel | None = None, + ) -> None: + self._provider = provider + self._credential_store = credential_store + self._thinking_level = thinking_level + + async def aclose(self) -> None: + """No persistent HTTP client is owned by the wrapper.""" + + def stream_response( + self, + *, + model: str, + system: str, + messages: list[AgentMessage], + tools: list[AgentTool], + signal: CancellationToken | None = None, + ) -> AsyncIterator[ProviderEvent]: + """Refresh Copilot credentials, then delegate one streamed request.""" + + async def iterator() -> AsyncIterator[ProviderEvent]: + inner = await self._fresh_inner_provider(model=model) + try: + async for event in inner.stream_response( + model=model, + system=system, + messages=messages, + tools=tools, + signal=signal, + ): + yield event + finally: + await inner.aclose() + + return iterator() + + async def _fresh_inner_provider(self, *, model: str) -> ClosableModelProvider: + provider = await self._fresh_provider_config() + return OpenAICompatibleProvider( + openai_compatible_config_from_provider( + provider, + credential_reader=self._credential_store, + model=model, + thinking_level=self._thinking_level, + ) + ) + + async def _fresh_provider_config(self) -> ProviderConfig: + headers = {**dict(self._provider.headers), **_github_copilot_headers()} + credential_name = self._provider.credential_name + if not credential_name: + return replace(self._provider, headers=headers) + + credential = self._credential_store.get_oauth(credential_name) + if credential is None: + return replace(self._provider, headers=headers) + if oauth_credential_is_expired(credential): + credential = await refresh_github_copilot_token( + credential.refresh, + enterprise_domain=credential.account_id if credential.account_id != "github.com" else "", + ) + self._credential_store.set_oauth(credential_name, credential) + return replace( + self._provider, + base_url=github_copilot_base_url(credential.access, credential.account_id), + headers=headers, + ) + + + +def _anthropic_provider_config_for_model( + provider: ProviderConfig, + model: str, +) -> AnthropicProviderConfig: + """Adapt shared connection settings for a model served via Messages API.""" + if isinstance(provider, AnthropicProviderConfig): + return provider + return AnthropicProviderConfig( + name=provider.name, + base_url=provider.base_url, + api_key_env=provider.api_key_env, + credential_name=provider.credential_name, + models=provider.models, + default_model=model, + context_windows=provider.context_windows, + headers=provider.headers, + timeout_seconds=provider.timeout_seconds, + max_retries=provider.max_retries, + max_retry_delay_seconds=provider.max_retry_delay_seconds, + thinking_levels=provider.thinking_levels, + thinking_models=provider.thinking_models, + thinking_default=provider.thinking_default, + thinking_parameter="anthropic.thinking", + ) + + def _codex_reasoning_effort( provider: OpenAICodexProviderConfig, *, diff --git a/src/tau_coding/session.py b/src/tau_coding/session.py index 110f0669a..ddb6690de 100644 --- a/src/tau_coding/session.py +++ b/src/tau_coding/session.py @@ -726,6 +726,11 @@ def set_provider(self, provider_name: str, *, persist_default: bool = True) -> N self._owned_providers.append(provider) self._harness.config.provider = provider self._provider_name = provider_config.name + self._diagnostic_logger.log_runtime_provider( + context=self._diagnostic_context(), + phase="set_provider_runtime", + provider=provider, + ) self._runtime_provider_config = provider_config self._harness.config.model = model self._thinking_level = thinking_level @@ -827,6 +832,11 @@ def _refresh_runtime_provider(self) -> None: self._owned_providers.append(provider) self._harness.config.provider = provider self._runtime_provider_config = provider_config + self._diagnostic_logger.log_runtime_provider( + context=self._diagnostic_context(), + phase="refresh_runtime_provider", + provider=provider, + ) def reload(self) -> CodingReloadSummary: """Reload local coding resources and project context for future turns.""" @@ -912,6 +922,7 @@ async def resume(self, session_id: str) -> str: provider_name = self._provider_name runtime_provider_config = self._runtime_provider_config + restored_model = record.model or self.model if record.provider_name: if self._provider_settings is None: raise ValueError( @@ -925,11 +936,20 @@ async def resume(self, session_id: str) -> str: f"Session provider is not configured: {record.provider_name}" ) from exc provider_name = runtime_provider_config.name + else: + inferred = _infer_provider_for_model( + self._provider_settings, + restored_model, + current_provider_name=self._provider_name, + ) + if inferred is not None: + provider_name = inferred.name + runtime_provider_config = inferred replacement = await type(self).load( CodingSessionConfig( provider=self._harness.config.provider, - model=record.model or self.model, + model=restored_model, cwd=record.cwd, storage=jsonl_session_storage(record.path), system=self._config.system, @@ -1733,6 +1753,34 @@ def _session_export_title(session: CodingSession) -> str: return f"Tau session {session_id}" if session_id is not None else "Tau Session Export" +def _infer_provider_for_model( + provider_settings: ProviderSettings | None, + model: str, + *, + current_provider_name: str, +) -> ProviderConfig | None: + """Infer the provider for a restored session model when index metadata is stale. + + Older session indexes did not always persist ``provider_name``. Resuming one + of those sessions after using another provider can otherwise combine the + restored model with the current runtime provider (for example + ``openai-codex:claude-opus-4.8``), which fails before the user can recover. + Only infer when provider settings identify a unique configured provider for + the model, and keep the current provider when it already advertises it. + """ + if provider_settings is None or not model: + return None + matches = tuple(provider for provider in provider_settings.providers if model in provider.models) + if not matches: + return None + for provider in matches: + if provider.name == current_provider_name: + return provider + if len(matches) == 1: + return matches[0] + return None + + def _state_thinking_level( state: SessionState, default: ThinkingLevel, diff --git a/src/tau_coding/session_export.py b/src/tau_coding/session_export.py index 5881860ed..4b554f400 100644 --- a/src/tau_coding/session_export.py +++ b/src/tau_coding/session_export.py @@ -536,7 +536,7 @@ def _render_message_entry(entry: MessageEntry) -> str: "{_escape(call.name)} "
f"{_escape(call.id)}"
- f"{_escape(_json_dump(call.arguments))}"
+ f"{_escape(_json_dump(_public_tool_arguments(call.arguments)))}"
"