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("") for call in message.tool_calls: - lines.append(f"- {call.name}: {call.arguments}") + arguments = _public_tool_arguments(call.arguments) + lines.append(f"- {call.name}: {arguments}") lines.append("") lines.append("") case "tool": @@ -250,6 +251,11 @@ def serialize_messages_for_compaction(messages: tuple[AgentMessage, ...]) -> str return "\n".join(lines) +def _public_tool_arguments(arguments: dict[str, object]) -> dict[str, object]: + return {key: value for key, value in arguments.items() if key != "_raw_arguments"} + + + def _message_text(message: AgentMessage) -> str: match message.role: case "user": diff --git a/src/tau_coding/diagnostics.py b/src/tau_coding/diagnostics.py index 8eec9ea35..3a9421ed9 100644 --- a/src/tau_coding/diagnostics.py +++ b/src/tau_coding/diagnostics.py @@ -53,6 +53,28 @@ def log_exception( self._append(entry) return self.path + def log_runtime_provider( + self, + *, + context: AgentCallDiagnosticContext, + phase: str, + provider: object, + ) -> Path: + """Log the concrete runtime provider object selected for a session phase.""" + entry = _base_entry(context, phase=phase, kind="runtime_provider") + provider_type = type(provider) + entry["runtime_provider"] = { + "class": provider_type.__qualname__, + "module": provider_type.__module__, + } + inner = getattr(provider, "_inner", None) + if inner is not None: + inner_type = type(inner) + entry["runtime_provider"]["inner_class"] = inner_type.__qualname__ + entry["runtime_provider"]["inner_module"] = inner_type.__module__ + self._append(entry) + return self.path + def log_error_event( self, *, diff --git a/src/tau_coding/oauth.py b/src/tau_coding/oauth.py index 6d3827a06..07d831fcb 100644 --- a/src/tau_coding/oauth.py +++ b/src/tau_coding/oauth.py @@ -21,6 +21,16 @@ from tau_coding.credentials import OAuthCredential +GITHUB_COPILOT_OAUTH_PROVIDER = "github-copilot" +GITHUB_COPILOT_CLIENT_ID = "Iv1.b507a08c87ecfe98" +GITHUB_COPILOT_API_VERSION = "2026-06-01" +GITHUB_COPILOT_HEADERS = { + "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_CODEX_OAUTH_PROVIDER = "openai-codex" OPENAI_CODEX_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann" OPENAI_CODEX_AUTHORIZE_URL = "https://auth.openai.com/oauth/authorize" @@ -501,3 +511,211 @@ def _oauth_html(message: str) -> str: .replace('"', """) ) return f'Tau OAuth

{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: "
  • " f"{_escape(call.name)} " f"{_escape(call.id)}" - f"
    {_escape(_json_dump(call.arguments))}
    " + f"
    {_escape(_json_dump(_public_tool_arguments(call.arguments)))}
    " "
  • " for call in message.tool_calls ) @@ -648,6 +648,11 @@ def _summarize_text(text: str, *, limit: int = 92) -> str: return summary[: limit - 3].rstrip() + "..." +def _public_tool_arguments(arguments: dict[str, JSONValue]) -> dict[str, JSONValue]: + return {key: value for key, value in arguments.items() if key != "_raw_arguments"} + + + def _json_dump(value: dict[str, JSONValue]) -> str: return json.dumps(value, indent=2, sort_keys=True) diff --git a/src/tau_coding/tui/app.py b/src/tau_coding/tui/app.py index 60968c648..e9979802f 100644 --- a/src/tau_coding/tui/app.py +++ b/src/tau_coding/tui/app.py @@ -56,7 +56,7 @@ from tau_ai.provider import CancellationToken from tau_coding.commands import CommandRegistry, create_default_command_registry from tau_coding.credentials import FileCredentialStore, OAuthCredential -from tau_coding.oauth import OAuthAuthInfo, OAuthPrompt, login_openai_codex +from tau_coding.oauth import OAuthAuthInfo, OAuthPrompt, login_github_copilot, login_openai_codex from tau_coding.provider_catalog import ( BUILTIN_PROVIDER_CATALOG, ProviderCatalogEntry, @@ -1340,28 +1340,48 @@ def __init__(self, provider: ProviderCatalogEntry, *, theme: TuiTheme) -> None: def compose(self) -> ComposeResult: """Compose the OAuth login prompt.""" + is_device_flow = self.provider.name == "github-copilot" with Vertical(id="login-screen"): yield Static(f"Login: {self.provider.display_name}", id="login-title") - yield Static("Complete the browser login, or paste the redirect URL.", id="login-help") + yield Static( + ( + "Open the URL and enter the displayed device code." + if is_device_flow + else "Complete the browser login, or paste the redirect URL." + ), + id="login-help", + ) yield Static("", id="login-oauth-url") yield Input( - placeholder="Paste redirect URL or authorization code", + placeholder=( + "Device login is automatic" + if is_device_flow + else "Paste redirect URL or authorization code" + ), id="login-oauth-code", + disabled=is_device_flow, + ) + yield Static( + "Escape closes" if is_device_flow else "Enter submits - Escape closes", + id="login-footer", ) - yield Static("Enter submits - Escape closes", id="login-footer") def on_mount(self) -> None: """Focus the manual-code field and start OAuth.""" - self.query_one("#login-oauth-code", Input).focus() + if self.provider.name != "github-copilot": + self.query_one("#login-oauth-code", Input).focus() self.run_worker(self._run_login(), exclusive=True) async def _run_login(self) -> None: try: - credential = await login_openai_codex( - on_auth=self._show_auth, - on_prompt=self._prompt_for_code, - on_manual_code_input=self._manual_code_input, - ) + if self.provider.name == "github-copilot": + credential = await login_github_copilot(on_auth=self._show_auth) + else: + credential = await login_openai_codex( + on_auth=self._show_auth, + on_prompt=self._prompt_for_code, + on_manual_code_input=self._manual_code_input, + ) except Exception as exc: # noqa: BLE001 - surface OAuth failures in the TUI self.query_one("#login-help", Static).update(f"OAuth failed: {exc}") return @@ -2626,7 +2646,7 @@ def _open_login(self, provider_name: str) -> None: if entry is None: self._notify(f"Unknown provider: {provider_name}", severity="error") return - if entry.kind == "openai-codex": + if entry.kind == "openai-codex" or entry.name == "github-copilot": self.push_screen( OAuthLoginScreen(entry, theme=self.tui_settings.resolved_theme), callback=lambda credential: self._handle_oauth_login_result(entry, credential), @@ -3367,13 +3387,21 @@ def _login_provider_label(provider: ProviderCatalogEntry) -> str: def _subscription_login_providers( providers: Sequence[ProviderCatalogEntry], ) -> tuple[ProviderCatalogEntry, ...]: - return tuple(provider for provider in providers if provider.kind == "openai-codex") + return tuple( + provider + for provider in providers + if provider.kind == "openai-codex" or provider.name == "github-copilot" + ) def _api_key_login_providers( providers: Sequence[ProviderCatalogEntry], ) -> tuple[ProviderCatalogEntry, ...]: - return tuple(provider for provider in providers if provider.kind != "openai-codex") + return tuple( + provider + for provider in providers + if provider.kind != "openai-codex" and provider.name != "github-copilot" + ) def _stored_credential_providers( diff --git a/src/tau_coding/tui/state.py b/src/tau_coding/tui/state.py index 4d1a98981..92a732ba0 100644 --- a/src/tau_coding/tui/state.py +++ b/src/tau_coding/tui/state.py @@ -275,11 +275,17 @@ def _read_line_suffix(arguments: dict[str, JSONValue]) -> str: def _fallback_tool_call_invocation(tool_call: ToolCall) -> str: - if tool_call.arguments: - return f"{tool_call.name} {tool_call.arguments}" + display_arguments = _display_tool_arguments(tool_call.arguments) + if display_arguments: + return f"{tool_call.name} {display_arguments}" return tool_call.name +def _display_tool_arguments(arguments: dict[str, JSONValue]) -> dict[str, JSONValue]: + """Return human-facing tool arguments without parser-internal fallback data.""" + return {key: value for key, value in arguments.items() if key != "_raw_arguments"} + + def _string_argument(arguments: dict[str, JSONValue], key: str) -> str | None: value = arguments.get(key) return value if isinstance(value, str) else None diff --git a/tests/test_agent_harness.py b/tests/test_agent_harness.py index 9d4360411..006df61ba 100644 --- a/tests/test_agent_harness.py +++ b/tests/test_agent_harness.py @@ -198,6 +198,60 @@ async def test_cancel_requests_cancellation_for_current_run() -> None: assert harness.messages == (UserMessage(content="Hi"),) +@pytest.mark.anyio +async def test_tool_result_id_rewrite_works_without_model_copy() -> None: + class LegacyToolResult: + tool_call_id = "wrong-id" + name = "read" + ok = True + content = "ok" + error = None + data = {"path": "README.md"} + + async def executor( + arguments: Mapping[str, JSONValue], + signal: object | None = None, + ) -> AgentToolResult: + del arguments, signal + return LegacyToolResult() # type: ignore[return-value] + + tool = AgentTool( + name="read", + description="Read a file.", + input_schema={"type": "object"}, + executor=executor, + ) + tool_call = ToolCall(id="call-1", name="read", arguments={"path": "README.md"}) + provider = FakeProvider( + [ + [ + ProviderResponseStartEvent(model="fake"), + ProviderResponseEndEvent( + message=AssistantMessage(content="I'll read it.", tool_calls=[tool_call]) + ), + ], + [ + ProviderResponseStartEvent(model="fake"), + ProviderResponseEndEvent(message=AssistantMessage(content="Done.")), + ], + ] + ) + harness = AgentHarness( + AgentHarnessConfig( + provider=provider, + model="fake", + system="You are Tau.", + tools=[tool], + ) + ) + + _events = [event async for event in harness.prompt("read")] + + result = next(message for message in harness.messages if isinstance(message, ToolResultMessage)) + assert result.tool_call_id == "call-1" + assert result.content == "ok" + + @pytest.mark.anyio async def test_cancelled_tool_run_repairs_transcript_before_next_prompt() -> None: async def executor( diff --git a/tests/test_coding_session.py b/tests/test_coding_session.py index 21ed769ba..6525f9326 100644 --- a/tests/test_coding_session.py +++ b/tests/test_coding_session.py @@ -2781,6 +2781,80 @@ def create_provider( assert created == [("local", "qwen")] +@pytest.mark.anyio +async def test_session_resume_infers_provider_for_legacy_index_without_provider_name( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) + first_record = manager.create_session( + cwd=tmp_path / "first", + model="gpt-5.5", + provider_name="openai-codex", + title="First", + ) + second_cwd = tmp_path / "second" + second_cwd.mkdir(parents=True) + second_record = manager.create_session( + cwd=second_cwd, + model="claude-opus-4.8", + provider_name=None, + title="Legacy Copilot Claude", + ) + assert second_record.provider_name is None + settings = ProviderSettings( + default_provider="openai-codex", + providers=( + OpenAICodexProviderConfig(models=("gpt-5.5",), default_model="gpt-5.5"), + OpenAICompatibleProviderConfig( + name="github-copilot", + base_url="https://api.githubcopilot.com", + api_key_env="GITHUB_COPILOT_TOKEN", + models=("claude-opus-4.8",), + default_model="claude-opus-4.8", + ), + ), + ) + created: list[tuple[str, str | None]] = [] + + def create_provider( + provider_config: object, + *, + credential_store: FileCredentialStore | None = None, + model: str | None = None, + thinking_level: str | None = None, + llm_observer: object | None = None, + ) -> SwitchableFakeProvider: + del credential_store, thinking_level, llm_observer + created.append((provider_config.name, model)) # type: ignore[attr-defined] + return SwitchableFakeProvider(provider_config) + + monkeypatch.setattr(coding_session_module, "create_model_provider", create_provider) + second_storage = JsonlSessionStorage(second_record.path) + await second_storage.append(SessionInfoEntry(cwd=str(second_record.cwd))) + await second_storage.append(ModelChangeEntry(model="claude-opus-4.8")) + session = await CodingSession.load( + CodingSessionConfig( + provider=FakeProvider([]), + model="gpt-5.5", + system="You are Tau.", + storage=JsonlSessionStorage(first_record.path), + cwd=first_record.cwd, + session_id=first_record.id, + session_manager=manager, + provider_name="openai-codex", + provider_settings=settings, + runtime_provider_config=settings.get_provider("openai-codex"), + ) + ) + created.clear() + + await session.resume(second_record.id) + + assert session.provider_name == "github-copilot" + assert session.model == "claude-opus-4.8" + assert created == [("github-copilot", "claude-opus-4.8")] + + @pytest.mark.anyio async def test_session_context_usage_recalculates_after_resume(tmp_path: Path) -> None: manager = SessionManager(TauPaths(home=tmp_path / ".tau", agents_home=tmp_path / ".agents")) @@ -2830,3 +2904,30 @@ def test_minimal_commands_are_handled(tmp_path: Path) -> None: assert session.handle_command("/quit").exit_requested is True assert session.handle_command("/exit").exit_requested is True assert session.handle_command("/unknown").message == "Unknown command: /unknown" + + +def test_branch_summary_hides_raw_argument_fallbacks() -> None: + from tau_coding.branch_summary import _format_assistant_summary_source + + summary = _format_assistant_summary_source( + AssistantMessage( + content="", + tool_calls=[ + ToolCall( + id="call-raw", + name="read", + arguments={"_raw_arguments": "{not valid json"}, + ), + ToolCall( + id="call-partial", + name="custom", + arguments={"path": "README.md", "_raw_arguments": "{not valid json"}, + ), + ], + ) + ) + + assert "_raw_arguments" not in summary + assert "{not valid json" not in summary + assert "read()" in summary + assert 'custom(path="README.md")' in summary diff --git a/tests/test_oauth.py b/tests/test_oauth.py index f683cef49..2a379b8a0 100644 --- a/tests/test_oauth.py +++ b/tests/test_oauth.py @@ -11,7 +11,10 @@ OPENAI_CODEX_CLIENT_ID, account_id_from_access_token, create_openai_codex_authorization_flow, + github_copilot_base_url, + normalize_github_domain, parse_authorization_input, + refresh_github_copilot_token, refresh_openai_codex_token, ) @@ -93,6 +96,49 @@ def handler(request: httpx.Request) -> httpx.Response: assert credential.expires == expires * 1000 +def test_github_copilot_base_url_extracts_proxy_endpoint_and_enterprise_domain() -> None: + token = "x;proxy-ep=proxy.enterprise.githubcopilot.com;y" + + assert ( + github_copilot_base_url(token) + == "https://api.enterprise.githubcopilot.com" + ) + assert ( + github_copilot_base_url("", "example.ghe.com") + == "https://copilot-api.example.ghe.com" + ) + assert ( + github_copilot_base_url("") + == "https://api.individual.githubcopilot.com" + ) + + +def test_normalize_github_domain_accepts_url_or_hostname() -> None: + assert normalize_github_domain("https://example.ghe.com/org") == "example.ghe.com" + assert normalize_github_domain("example.ghe.com") == "example.ghe.com" + assert normalize_github_domain("") == "" + + +@pytest.mark.anyio +async def test_refresh_github_copilot_token_returns_copilot_oauth_credential() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert request.url == "https://api.github.com/copilot_internal/v2/token" + assert request.headers["authorization"] == "Bearer github-token" + assert request.headers["copilot-integration-id"] == "vscode-chat" + return httpx.Response( + 200, + json={"token": "copilot-token", "expires_at": 2_000_000}, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + credential = await refresh_github_copilot_token("github-token", client=client) + + assert credential.access == "copilot-token" + assert credential.refresh == "github-token" + assert credential.account_id == "github.com" + assert credential.expires == 2_000_000_000 - 5 * 60 * 1000 + + def _jwt(account_id: str, *, expires: int | None = None) -> str: payload = {OPENAI_CODEX_ACCOUNT_CLAIM: {"chatgpt_account_id": account_id}} if expires is not None: diff --git a/tests/test_provider_config.py b/tests/test_provider_config.py index 845cf07f0..b93ab3eac 100644 --- a/tests/test_provider_config.py +++ b/tests/test_provider_config.py @@ -1,6 +1,7 @@ import json from pathlib import Path +import httpx import pytest from tau_coding.credentials import FileCredentialStore, OAuthCredential @@ -14,6 +15,7 @@ ProviderSettings, ScopedModelConfig, anthropic_config_from_provider, + ensure_dynamic_provider_models, load_provider_settings, openai_compatible_config_from_provider, provider_default_thinking_level, @@ -35,8 +37,12 @@ def test_load_provider_settings_missing_file_uses_openai_default(tmp_path: Path) "openai", "openai-codex", "anthropic", + "github-copilot", "openrouter", "huggingface", + "deepseek", + "opencode-go", + "nebius", ] assert settings.providers[0].default_model == DEFAULT_MODEL assert settings.get_provider("anthropic").api_key_env == "ANTHROPIC_API_KEY" @@ -51,6 +57,7 @@ def test_builtin_openai_declares_model_scoped_thinking_capabilities() -> None: huggingface = settings.get_provider("huggingface") codex = settings.get_provider("openai-codex") anthropic = settings.get_provider("anthropic") + copilot = settings.get_provider("github-copilot") assert openai.context_windows["gpt-5.5"] == 272_000 assert openai.context_windows["gpt-5.5-pro"] == 1_050_000 @@ -111,6 +118,19 @@ def test_builtin_openai_declares_model_scoped_thinking_capabilities() -> None: ) assert provider_thinking_unavailable_reason(anthropic, model="claude-sonnet-4-6") is None assert provider_thinking_levels(anthropic, model="claude-haiku-4-5") == () + assert copilot.default_model == "gpt-5.5" + assert copilot.dynamic_models is True + assert "claude-sonnet-5" in copilot.models + assert "claude-sonnet-4.6" in copilot.models + assert "gemini-3.5-flash" in copilot.models + assert provider_thinking_levels(copilot, model="gpt-5.4") == ( + "off", + "low", + "medium", + "high", + "xhigh", + ) + assert provider_thinking_levels(copilot, model="claude-sonnet-4.6") == () def test_save_provider_settings_writes_backup_when_replacing(tmp_path: Path) -> None: @@ -148,7 +168,12 @@ def test_save_provider_settings_writes_backup_when_replacing(tmp_path: Path) -> ) -def test_save_and_load_provider_settings_round_trip(tmp_path: Path) -> None: +def test_save_and_load_provider_settings_round_trip( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + monkeypatch.delenv("DEEPSEEK_API_KEY", raising=False) + monkeypatch.delenv("OPENROUTER_API_KEY", raising=False) paths = TauPaths(home=tmp_path / ".tau") settings = ProviderSettings( default_provider="local", @@ -234,10 +259,14 @@ def test_upsert_openai_compatible_provider_replaces_and_sets_default() -> None: assert updated.default_provider == "local" assert [item.name for item in updated.providers] == [ "anthropic", + "deepseek", + "github-copilot", "huggingface", "local", + "nebius", "openai", "openai-codex", + "opencode-go", "openrouter", ] assert replaced.get_provider("local").default_model == "llama" @@ -694,6 +723,7 @@ def test_load_provider_settings_restores_builtin_providers_with_stored_credentia tmp_path: Path, ) -> None: for env_name in ( + "DEEPSEEK_API_KEY", "OPENAI_API_KEY", "OPENAI_CODEX_ACCESS_TOKEN", "ANTHROPIC_API_KEY", @@ -860,3 +890,166 @@ def test_openai_compatible_provider_config_rejects_invalid_retries() -> None: OpenAICompatibleProviderConfig(name="local", max_retries=-1) with pytest.raises(ProviderConfigError, match="0 or greater"): OpenAICompatibleProviderConfig(name="local", max_retry_delay_seconds=-1) + + +def test_nebius_builtin_entry_is_dynamic_with_empty_catalog() -> None: + settings = ProviderSettings() + nebius = settings.get_provider("nebius") + + assert isinstance(nebius, OpenAICompatibleProviderConfig) + assert nebius.base_url == "https://api.tokenfactory.nebius.com/v1" + assert nebius.api_key_env == "NEBIUS_TOKEN_FACTORY_API_KEY" + assert nebius.models == () + assert nebius.default_model == "" + + +@pytest.mark.anyio +async def test_ensure_dynamic_provider_models_populates_nebius_at_build( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("NEBIUS_TOKEN_FACTORY_API_KEY", "nebius-key") + paths = TauPaths(home=tmp_path / ".tau") + settings = load_provider_settings(paths) + + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.params["verbose"] == "true" + return httpx.Response( + 200, + json={ + "data": [ + {"id": "meta-llama/Llama-3.3-70B-Instruct", "context_window": 131072}, + {"id": "deepseek-ai/DeepSeek-R1-0528"}, + ] + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + updated = await ensure_dynamic_provider_models( + settings, provider_name="nebius", paths=paths, client=client + ) + + nebius = updated.get_provider("nebius") + assert nebius.models == ( + "meta-llama/Llama-3.3-70B-Instruct", + "deepseek-ai/DeepSeek-R1-0528", + ) + assert nebius.default_model == "meta-llama/Llama-3.3-70B-Instruct" + assert nebius.context_windows["meta-llama/Llama-3.3-70B-Instruct"] == 131072 + + reloaded = load_provider_settings(paths) + assert reloaded.get_provider("nebius").models == nebius.models + + +@pytest.mark.anyio +async def test_ensure_dynamic_provider_models_uses_copilot_proxy_and_headers( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("GITHUB_COPILOT_TOKEN", raising=False) + paths = TauPaths(home=tmp_path / ".tau") + credential_store = FileCredentialStore(tmp_path / ".tau" / "credentials.json") + credential_store.set_oauth( + "github-copilot", + OAuthCredential( + access="tid=1;proxy-ep=proxy.enterprise.test;token", + refresh="github-refresh", + expires=9999999999999, + account_id="github.com", + ), + ) + settings = load_provider_settings(paths) + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + json={ + "data": [ + {"id": "gpt-5.5", "context_window": 272000}, + {"id": "claude-sonnet-5", "context_window": 1000000}, + {"id": "gemini-3.5-flash"}, + ] + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + updated = await ensure_dynamic_provider_models( + settings, + provider_name="github-copilot", + paths=paths, + credential_store=credential_store, + client=client, + ) + + copilot = updated.get_provider("github-copilot") + assert requests[0].url == "https://api.enterprise.test/models?verbose=true" + assert requests[0].headers["authorization"] == "Bearer tid=1;proxy-ep=proxy.enterprise.test;token" + assert requests[0].headers["copilot-integration-id"] == "vscode-chat" + assert copilot.models == ("gpt-5.5", "claude-sonnet-5", "gemini-3.5-flash") + assert copilot.default_model == "gpt-5.5" + assert copilot.context_windows["claude-sonnet-5"] == 1000000 + + +@pytest.mark.anyio +async def test_ensure_dynamic_provider_models_leaves_non_dynamic_unchanged( + tmp_path: Path, +) -> None: + paths = TauPaths(home=tmp_path / ".tau") + settings = load_provider_settings(paths) + openai_before = settings.get_provider("openai") + + updated = await ensure_dynamic_provider_models( + settings, provider_name="openai", paths=paths + ) + + assert updated.get_provider("openai").models == openai_before.models + + +@pytest.mark.anyio +async def test_ensure_dynamic_provider_models_without_credentials_is_noop( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("NEBIUS_TOKEN_FACTORY_API_KEY", raising=False) + paths = TauPaths(home=tmp_path / ".tau") + settings = load_provider_settings(paths) + + def handler(request: httpx.Request) -> httpx.Response: + raise AssertionError("should not call the models endpoint without credentials") + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + updated = await ensure_dynamic_provider_models( + settings, provider_name="nebius", paths=paths, client=client + ) + + assert updated.get_provider("nebius").models == () + + +def test_github_copilot_provider_uses_stored_oauth_access_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("GITHUB_COPILOT_TOKEN", raising=False) + provider = ProviderSettings().get_provider("github-copilot") + + class FakeCredentials: + def get(self, name: str) -> str | None: + return None + + def get_oauth(self, name: str) -> OAuthCredential | None: + if name != "github-copilot": + return None + return OAuthCredential( + access="copilot-token", + refresh="github-token", + expires=999_999_999_999, + account_id="github.com", + ) + + config = openai_compatible_config_from_provider( + provider, + credential_reader=FakeCredentials(), + ) + + assert config.api_key == "copilot-token" diff --git a/tests/test_provider_runtime.py b/tests/test_provider_runtime.py index 57681dbd2..3f6ed50af 100644 --- a/tests/test_provider_runtime.py +++ b/tests/test_provider_runtime.py @@ -1,10 +1,14 @@ import pytest -from tau_ai import OpenAICodexProvider +from tau_ai import OpenAICodexProvider, OpenAICompatibleProvider from tau_coding import provider_runtime from tau_coding.credentials import FileCredentialStore, OAuthCredential -from tau_coding.provider_config import OpenAICodexProviderConfig -from tau_coding.provider_runtime import OpenAICodexCredentialResolver, create_model_provider +from tau_coding.provider_config import OpenAICodexProviderConfig, ProviderSettings +from tau_coding.provider_runtime import ( + GitHubCopilotCredentialRefreshingProvider, + OpenAICodexCredentialResolver, + create_model_provider, +) def test_create_model_provider_returns_openai_codex_provider(tmp_path) -> None: @@ -18,6 +22,43 @@ def test_create_model_provider_returns_openai_codex_provider(tmp_path) -> None: assert isinstance(provider, OpenAICodexProvider) +@pytest.mark.anyio +@pytest.mark.parametrize( + "model", + ["gpt-5.5", "claude-sonnet-5", "gemini-3.5-flash"], +) +async def test_create_model_provider_keeps_github_copilot_models_openai_compatible( + tmp_path, + model: str, +) -> None: + store = FileCredentialStore(tmp_path / "credentials.json") + store.set_oauth( + "github-copilot", + OAuthCredential( + access="tid=1;proxy-ep=proxy.enterprise.test;token", + refresh="github-refresh", + expires=9999999999999, + account_id="github.com", + ), + ) + provider_config = ProviderSettings().get_provider("github-copilot") + + provider = create_model_provider( + provider_config, + credential_store=store, + model=model, + ) + + assert isinstance(provider, GitHubCopilotCredentialRefreshingProvider) + inner = await provider._fresh_inner_provider(model=model) + try: + assert isinstance(inner, OpenAICompatibleProvider) + assert inner._config.base_url == "https://api.enterprise.test" + assert inner._config.headers["Copilot-Integration-Id"] == "vscode-chat" + finally: + await inner.aclose() + + def test_create_model_provider_maps_codex_reasoning_effort_like_pi(tmp_path) -> None: store = FileCredentialStore(tmp_path / "credentials.json") provider_config = OpenAICodexProviderConfig( @@ -53,6 +94,57 @@ def test_create_model_provider_maps_codex_reasoning_effort_like_pi(tmp_path) -> assert xhigh_provider._config.reasoning_effort == "xhigh" +@pytest.mark.anyio +async def test_github_copilot_provider_refreshes_expired_credentials_per_request( + monkeypatch: pytest.MonkeyPatch, + tmp_path, +) -> None: + store = FileCredentialStore(tmp_path / "credentials.json") + store.set_oauth( + "github-copilot", + OAuthCredential( + access="old-token", + refresh="github-refresh", + expires=1, + account_id="github.com", + ), + ) + provider_config = ProviderSettings().get_provider("github-copilot") + + async def fake_refresh(refresh_token: str, *, enterprise_domain: str = "") -> OAuthCredential: + assert refresh_token == "github-refresh" + assert enterprise_domain == "" + return OAuthCredential( + access="tid=1;proxy-ep=proxy.enterprise.test;new-token", + refresh="github-refresh", + expires=9999999999999, + account_id="github.com", + ) + + monkeypatch.setattr(provider_runtime, "refresh_github_copilot_token", fake_refresh) + + provider = create_model_provider( + provider_config, + credential_store=store, + model="claude-sonnet-5", + ) + + assert isinstance(provider, GitHubCopilotCredentialRefreshingProvider) + inner = await provider._fresh_inner_provider(model="claude-sonnet-5") + try: + assert isinstance(inner, OpenAICompatibleProvider) + assert inner._config.api_key == "tid=1;proxy-ep=proxy.enterprise.test;new-token" + assert inner._config.base_url == "https://api.enterprise.test" + finally: + await inner.aclose() + assert store.get_oauth("github-copilot") == OAuthCredential( + access="tid=1;proxy-ep=proxy.enterprise.test;new-token", + refresh="github-refresh", + expires=9999999999999, + account_id="github.com", + ) + + @pytest.mark.anyio async def test_openai_codex_credential_resolver_refreshes_expired_credentials( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_tau_ai.py b/tests/test_tau_ai.py index 84f3c8d61..7d24119aa 100644 --- a/tests/test_tau_ai.py +++ b/tests/test_tau_ai.py @@ -18,6 +18,7 @@ AnthropicConfig, AnthropicProvider, FakeProvider, + ModelInfo, OpenAICodexConfig, OpenAICodexCredentials, OpenAICodexProvider, @@ -30,6 +31,7 @@ ProviderTextDeltaEvent, ProviderThinkingDeltaEvent, ProviderToolCallEvent, + list_openai_compatible_models, openai_compatible_config_from_env, ) @@ -1538,3 +1540,180 @@ def handler(_request: httpx.Request) -> httpx.Response: assert isinstance(end, ProviderResponseEndEvent) assert end.message.content == "partial" assert end.finish_reason == "length" +@pytest.mark.anyio +async def test_list_openai_compatible_models_uses_verbose_and_parses_ids() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + json={ + "object": "list", + "data": [ + { + "id": "meta-llama/Llama-3.3-70B-Instruct", + "object": "model", + "owned_by": "system", + "context_window": 131072, + }, + {"id": "deepseek-ai/DeepSeek-R1-0528", "object": "model"}, + {"id": "meta-llama/Llama-3.3-70B-Instruct"}, + ], + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + models = await list_openai_compatible_models( + OpenAICompatibleConfig( + api_key="nebius-key", + base_url="https://api.tokenfactory.nebius.com/v1", + ), + verbose=True, + client=client, + ) + + assert [model.id for model in models] == [ + "meta-llama/Llama-3.3-70B-Instruct", + "deepseek-ai/DeepSeek-R1-0528", + ] + assert isinstance(models[0], ModelInfo) + assert models[0].context_window == 131072 + assert models[1].context_window is None + assert len(requests) == 1 + assert requests[0].url.path == "/v1/models" + assert requests[0].url.params["verbose"] == "true" + assert requests[0].headers["authorization"] == "Bearer nebius-key" + + +@pytest.mark.anyio +async def test_list_openai_compatible_models_omits_verbose_when_false() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert "verbose" not in request.url.params + return httpx.Response(200, json={"data": [{"id": "model-a"}]}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + models = await list_openai_compatible_models( + OpenAICompatibleConfig(api_key="key", base_url="https://example.test/v1"), + verbose=False, + client=client, + ) + + assert [model.id for model in models] == ["model-a"] + + + + +def test_openai_responses_payload_sanitizes_foreign_tool_ids() -> None: + from tau_ai.openai_compatible import _messages_to_responses_input + + foreign_id = "call_" + ("x" * 90) + "|fc.bad/id:with:separators" + items = _messages_to_responses_input( + [ + AssistantMessage( + content="", + tool_calls=[ToolCall(id=foreign_id, name="read", arguments={"path": "x"})], + ), + ToolResultMessage(tool_call_id=foreign_id, name="read", content="ok", ok=True), + ] + ) + + assert items[0]["type"] == "function_call" + assert items[0]["name"] == "read" + assert isinstance(items[0]["call_id"], str) + assert len(items[0]["call_id"]) <= 64 + assert "|" not in items[0]["call_id"] + assert items[1]["call_id"] == items[0]["call_id"] + + +def test_openai_codex_payload_sanitizes_foreign_tool_ids_and_empty_names() -> None: + from tau_ai.openai_codex import _build_codex_payload + + foreign_id = "call_" + ("x" * 90) + "|fc.bad/id:with:separators" + payload = _build_codex_payload( + model="gpt-5.5", + system="system", + messages=[ + AssistantMessage( + content="", + tool_calls=[ToolCall(id=foreign_id, name="", arguments={"path": "x"})], + ), + ToolResultMessage(tool_call_id=foreign_id, name="", content="ok", ok=True), + ], + tools=[], + ) + items = payload["input"] + + assert items[0]["type"] == "function_call" + assert items[0]["name"] == "tool" + assert isinstance(items[0]["call_id"], str) + assert len(items[0]["call_id"]) <= 64 + assert "|" not in items[0]["call_id"] + assert items[1]["call_id"] == items[0]["call_id"] + + +def test_anthropic_messages_payload_sanitizes_foreign_tool_ids() -> None: + from tau_ai.anthropic import _build_messages_payload + + foreign_id = "call_" + ("x" * 90) + "|fc.bad/id:with:separators" + payload = _build_messages_payload( + model="claude-test", + system="system", + messages=[ + AssistantMessage( + content="", + tool_calls=[ToolCall(id=foreign_id, name="read", arguments={"path": "x"})], + ), + ToolResultMessage(tool_call_id=foreign_id, name="read", content="ok", ok=True), + ], + tools=[], + thinking_budget_tokens=None, + ) + + tool_use_id = payload["messages"][0]["content"][0]["id"] + tool_result_id = payload["messages"][1]["content"][0]["tool_use_id"] + assert tool_result_id == tool_use_id + assert len(tool_use_id) <= 64 + assert "|" not in tool_use_id + assert "." not in tool_use_id + assert "/" not in tool_use_id + assert ":" not in tool_use_id + + +@pytest.mark.anyio +async def test_openai_chat_completions_replays_empty_tool_names_as_tool() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + text='data: {"choices":[{"delta":{"content":"ok"},"finish_reason":"stop"}]}\n\n', + headers={"content-type": "text/event-stream"}, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + provider = OpenAICompatibleProvider( + OpenAICompatibleConfig(api_key="test-key", base_url="https://example.test/v1"), + client=client, + ) + await _collect( + provider.stream_response( + model="gpt-5.1", + system="system", + messages=[ + AssistantMessage( + content="", + tool_calls=[ToolCall(id="call_1", name="", arguments={})], + ), + ToolResultMessage(tool_call_id="call_1", name="", content="ok", ok=True), + ], + tools=[], + ) + ) + + payload = loads(requests[0].content) + assistant = next(message for message in payload["messages"] if message["role"] == "assistant") + tool = next(message for message in payload["messages"] if message["role"] == "tool") + assert assistant["tool_calls"][0]["function"]["name"] == "tool" + assert tool["name"] == "tool" diff --git a/tests/test_tui_adapter.py b/tests/test_tui_adapter.py index d81179656..c25e72a1b 100644 --- a/tests/test_tui_adapter.py +++ b/tests/test_tui_adapter.py @@ -204,6 +204,29 @@ def test_tool_call_blocks_use_human_readable_invocations() -> None: ) +def test_tool_call_blocks_hide_raw_argument_fallbacks() -> None: + assert ( + format_tool_call_block( + ToolCall( + id="call-raw", + name="read", + arguments={"_raw_arguments": "{not valid json"}, + ) + ) + == "→ read" + ) + assert ( + format_tool_call_block( + ToolCall( + id="call-partial", + name="custom", + arguments={"path": "README.md", "_raw_arguments": "{not valid json"}, + ) + ) + == "→ custom {'path': 'README.md'}" + ) + + def test_tui_adapter_records_tool_updates_and_results() -> None: state = TuiState() adapter = TuiEventAdapter(state) diff --git a/tests/test_tui_app.py b/tests/test_tui_app.py index b5bb8ffea..741930058 100644 --- a/tests/test_tui_app.py +++ b/tests/test_tui_app.py @@ -3728,7 +3728,10 @@ async def test_tui_login_subscription_opens_oauth_provider_picker() -> None: assert isinstance(app.screen, LoginProviderPickerScreen) provider_list = app.screen.query_one("#login-provider-list", ListView) labels = [str(item.query_one(Label).render()) for item in provider_list.children] - assert labels == ["OpenAI Codex subscription\n openai-codex"] + assert labels == [ + "OpenAI Codex subscription\n openai-codex", + "GitHub Copilot\n github-copilot", + ] assert "gpt-5.5" not in "\n".join(labels) @@ -3752,6 +3755,7 @@ async def test_tui_login_api_key_opens_api_provider_picker() -> None: labels = [str(item.query_one(Label).render()) for item in provider_list.children] assert labels[0] == "OpenAI\n openai" assert "OpenAI Codex subscription\n openai-codex" not in labels + assert "GitHub Copilot\n github-copilot" not in labels await pilot.press("down") await pilot.press("enter")