diff --git a/__init__.py b/__init__.py index 38b110e..f05fadc 100644 --- a/__init__.py +++ b/__init__.py @@ -34,7 +34,7 @@ from providers.base import ProviderProfile # Override at runtime with the CLAUDE_CODE_CLI_BASE_URL env var. _DEFAULT_SHIM_BASE_URL = "http://127.0.0.1:8765/v1" -claude_code_cli = ProviderProfile( +_profile_kwargs = dict( name="claude-code-cli", # Aliases deliberately avoid "claude-code" - that one belongs to the native # `anthropic` profile. These point only at this local-CLI provider. @@ -51,13 +51,20 @@ claude_code_cli = ProviderProfile( env_vars=("CLAUDE_CODE_CLI_API_KEY", "CLAUDE_CODE_CLI_BASE_URL"), base_url=_DEFAULT_SHIM_BASE_URL, auth_type="api_key", - supports_vision=True, # Model ids the shim accepts and forwards verbatim to `claude --model`. # Shown in pickers when the live /models probe is unavailable. fallback_models=("opus", "sonnet", "haiku"), default_aux_model="haiku", ) +# Hermes v0.14 predates the declarative ``supports_vision`` profile field. +# Keep one plugin tree compatible with that production version and newer +# releases while the local shim continues to accept OpenAI image parts. +if "supports_vision" in getattr(ProviderProfile, "__dataclass_fields__", {}): + _profile_kwargs["supports_vision"] = True + +claude_code_cli = ProviderProfile(**_profile_kwargs) + register_provider(claude_code_cli) diff --git a/claude_code_server.py b/claude_code_server.py index 795d4bc..29606cc 100755 --- a/claude_code_server.py +++ b/claude_code_server.py @@ -83,13 +83,28 @@ def _env(name: str, default: str = "") -> str: def build_subprocess_env() -> dict[str, str]: - """Return an env with user-local CLI paths restored. + """Return a minimal env with user-local CLI paths restored. Mirrors fusion_runner.build_subprocess_env: cron/gateway-launched parents can have a minimal PATH that lacks ~/.local/npm/bin and nvm's node, which breaks `claude`'s `/usr/bin/env node` shebang even when the binary resolves. + The shim may run under Hermes, whose process environment contains provider + and channel secrets. Pass only the OS/session values Claude Code needs plus + its explicit subscription credential knobs. """ - env = os.environ.copy() + allowed = { + "HOME", "USER", "LOGNAME", "LANG", "LANGUAGE", "LC_ALL", "TERM", + "COLORTERM", "TMPDIR", "TMP", "TEMP", "XDG_CONFIG_HOME", + "XDG_CACHE_HOME", "XDG_DATA_HOME", "XDG_RUNTIME_DIR", + "SSL_CERT_FILE", "SSL_CERT_DIR", "NODE_EXTRA_CA_CERTS", + "CLAUDE_CODE_OAUTH_TOKEN", "CLAUDE_CONFIG_DIR", + "DISABLE_TELEMETRY", "DISABLE_ERROR_REPORTING", "DISABLE_AUTOUPDATER", + } + env = { + key: value + for key, value in os.environ.items() + if key in allowed or key.startswith("LC_") + } home = pathlib.Path.home() candidates: list[pathlib.Path] = [home / ".local/npm/bin", home / ".local/bin"] @@ -744,13 +759,13 @@ def run_claude(prompt: str, model: str, engine: bool = False, overrides: dict | if parsed is not None: if parsed.get("is_error"): msg = str(parsed.get("result") or parsed.get("error") or "claude reported is_error") - return {"text": f"[claude-code-cli error] {msg}", "usage": _map_usage(parsed.get("usage")), "error": msg} + return {"text": f"[claude-code-cli error] {msg}", "usage": _map_usage(parsed.get("usage"), prompt), "error": msg} served = _extract_served_model(parsed, model) substitution = check_model_provenance(model, served) if substitution: return { "text": f"[claude-code-cli error] {substitution}", - "usage": _map_usage(parsed.get("usage")), + "usage": _map_usage(parsed.get("usage"), prompt), "error": substitution, "model": served, } @@ -760,7 +775,7 @@ def run_claude(prompt: str, model: str, engine: bool = False, overrides: dict | if isinstance(result, str) and result.strip(): return { "text": result.strip(), - "usage": _map_usage(parsed.get("usage")), + "usage": _map_usage(parsed.get("usage"), prompt), "error": "", "model": served, } @@ -867,17 +882,21 @@ def _zero_usage() -> dict: return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0} -def _map_usage(usage) -> dict: - if not isinstance(usage, dict): - return _zero_usage() - prompt = ( - int(usage.get("input_tokens", 0) or 0) - + int(usage.get("cache_read_input_tokens", 0) or 0) - + int(usage.get("cache_creation_input_tokens", 0) or 0) - ) - completion = int(usage.get("output_tokens", 0) or 0) - return {"prompt_tokens": prompt, "completion_tokens": completion, - "total_tokens": prompt + completion} +def _estimate_prompt_tokens(prompt: str) -> int: + """Estimate the OpenAI request size for Hermes context accounting. + + Claude Code usage includes its own large system prompt and cache traffic, + which are outside the host request and can make Hermes compact too early. + A stable chars/4 estimate is preferable to reporting that internal usage. + """ + return (len(prompt) + 3) // 4 if prompt else 0 + + +def _map_usage(usage, prompt: str = "") -> dict: + prompt_tokens = _estimate_prompt_tokens(prompt) + completion = int(usage.get("output_tokens", 0) or 0) if isinstance(usage, dict) else 0 + return {"prompt_tokens": prompt_tokens, "completion_tokens": completion, + "total_tokens": prompt_tokens + completion} # --------------------------------------------------------------------------- # @@ -962,7 +981,7 @@ def _extract_assistant_message_text(event: dict) -> str: return "" -def _extract_result_event(event: dict): +def _extract_result_event(event: dict, prompt: str = ""): """Return ``(text, usage, error)`` for a terminal event, else ``None``. Recognizes the stream-json ``{"type":"result",...}`` event *and*, defensively, @@ -976,10 +995,10 @@ def _extract_result_event(event: dict): return None if event.get("is_error"): msg = str(event.get("result") or event.get("error") or "claude reported is_error") - return ("", _map_usage(event.get("usage")), msg) + return ("", _map_usage(event.get("usage"), prompt), msg) text = event.get("result") text = text if isinstance(text, str) else str(event.get("content") or "") - return (text, _map_usage(event.get("usage")), "") + return (text, _map_usage(event.get("usage"), prompt), "") def _kill_proc(proc) -> None: @@ -1098,7 +1117,7 @@ def stream_claude(prompt: str, model: str, engine: bool = False, yield ("delta", assistant_text) else: served_model = _extract_served_model(event, model) or served_model - parsed = _extract_result_event(event) + parsed = _extract_result_event(event, prompt) if parsed is not None: result_text, result_usage, result_error = parsed saw_result = True diff --git a/tests/test_claude_code_server.py b/tests/test_claude_code_server.py index 8aef2da..41c25cd 100644 --- a/tests/test_claude_code_server.py +++ b/tests/test_claude_code_server.py @@ -44,6 +44,54 @@ def make_fake_claude(directory: pathlib.Path, body: str) -> pathlib.Path: class ClaudeCodeServerUnitTests(unittest.TestCase): + def test_build_subprocess_env_keeps_subscription_auth_and_drops_provider_secrets(self): + with temporary_env( + ANTHROPIC_API_KEY="metered-api-key", + ANTHROPIC_AUTH_TOKEN="metered-auth-token", + ANTHROPIC_BASE_URL="https://metered.invalid", + CLAUDE_CODE_USE_BEDROCK="1", + CLAUDE_CODE_USE_VERTEX="1", + CLAUDE_CODE_USE_FOUNDRY="1", + OPENAI_API_KEY="other-provider-secret", + TELEGRAM_BOT_TOKEN="telegram-secret", + CLAUDE_CODE_OAUTH_TOKEN="subscription-oauth", + ): + child = server.build_subprocess_env() + + forbidden = { + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_BASE_URL", + "CLAUDE_CODE_USE_BEDROCK", + "CLAUDE_CODE_USE_VERTEX", + "CLAUDE_CODE_USE_FOUNDRY", + "OPENAI_API_KEY", + "TELEGRAM_BOT_TOKEN", + } + leaked = sorted(forbidden.intersection(child)) + self.assertEqual(leaked, []) + self.assertEqual(child["CLAUDE_CODE_OAUTH_TOKEN"], "subscription-oauth") + self.assertTrue(child["PATH"]) + + def test_map_usage_estimates_only_the_host_prompt(self): + usage = { + "input_tokens": 100, + "cache_read_input_tokens": 5000, + "cache_creation_input_tokens": 5000, + "output_tokens": 10, + } + + short = server._map_usage(usage, "hello") + long = server._map_usage(usage, "hello" + ("x" * 4000)) + + self.assertEqual(short, { + "prompt_tokens": 2, + "completion_tokens": 10, + "total_tokens": 12, + }) + self.assertEqual(long["prompt_tokens"], 1002) + self.assertGreater(long["prompt_tokens"], short["prompt_tokens"]) + def test_flatten_messages_preserves_system_conversation_and_tool_context(self): prompt = server.flatten_messages([ {"role": "system", "content": "System rule."}, @@ -134,7 +182,7 @@ print(json.dumps({ outcome = server.run_claude("hello world", "haiku", engine=False) self.assertEqual(outcome["text"], "fake response to world") - self.assertEqual(outcome["usage"], {"prompt_tokens": 10, "completion_tokens": 7, "total_tokens": 17}) + self.assertEqual(outcome["usage"], {"prompt_tokens": 3, "completion_tokens": 7, "total_tokens": 10}) self.assertEqual(outcome["error"], "") def test_run_claude_surfaces_nonzero_exit_without_crashing(self): @@ -179,7 +227,14 @@ print(json.dumps({"result": "hello from fake http", "usage": {"input_tokens": 1, completion = json.loads(request.urlopen(req, timeout=5).read().decode("utf-8")) self.assertEqual(completion["object"], "chat.completion") self.assertEqual(completion["choices"][0]["message"]["content"], "hello from fake http") - self.assertEqual(completion["usage"]["total_tokens"], 3) + expected_prompt = server._estimate_prompt_tokens(server.flatten_messages([ + {"role": "user", "content": "say hi"}, + ])) + self.assertEqual(completion["usage"], { + "prompt_tokens": expected_prompt, + "completion_tokens": 2, + "total_tokens": expected_prompt + 2, + }) stream_req = request.Request( base + "/v1/chat/completions", @@ -300,7 +355,13 @@ print(json.dumps({ self.assertEqual(calls[0]["type"], "function") self.assertEqual(calls[0]["function"]["name"], "skill_view") self.assertEqual(json.loads(calls[0]["function"]["arguments"]), {"name": "hermes-agent"}) - self.assertEqual(completion["usage"]["total_tokens"], 9) + self.assertEqual(completion["usage"]["completion_tokens"], 5) + self.assertGreater(completion["usage"]["prompt_tokens"], 0) + self.assertLess(completion["usage"]["prompt_tokens"], 1000) + self.assertEqual( + completion["usage"]["total_tokens"], + completion["usage"]["prompt_tokens"] + 5, + ) def test_native_tool_call_fallbacks_to_text_on_non_json_result(self): with tempfile.TemporaryDirectory() as td: @@ -776,7 +837,7 @@ emit({"type": "result", "result": "hello", "usage": {"input_tokens": 4, "output_ self.assertEqual(len(finals), 1) final = finals[0] self.assertEqual(final["text"], "hello") - self.assertEqual(final["usage"], {"prompt_tokens": 4, "completion_tokens": 5, "total_tokens": 9}) + self.assertEqual(final["usage"], {"prompt_tokens": 2, "completion_tokens": 5, "total_tokens": 7}) def test_stream_claude_bare_json_result_falls_back_to_single_delta(self): with tempfile.TemporaryDirectory() as td: @@ -838,7 +899,14 @@ emit({"type": "result", "result": "onetwo", "usage": {"input_tokens": 1, "output thread.join(timeout=5) self.assertIn('"content": "one"', body) self.assertIn('"content": "two"', body) - self.assertIn('"usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3}', body) + expected_prompt = server._estimate_prompt_tokens(server.flatten_messages([ + {"role": "user", "content": "say hi"}, + ])) + self.assertIn( + f'"usage": {{"prompt_tokens": {expected_prompt}, "completion_tokens": 2, ' + f'"total_tokens": {expected_prompt + 2}}}', + body, + ) self.assertIn("data: [DONE]", body) diff --git a/tests/test_provider_plugin_registration.py b/tests/test_provider_plugin_registration.py index f9c6165..bec256f 100644 --- a/tests/test_provider_plugin_registration.py +++ b/tests/test_provider_plugin_registration.py @@ -41,7 +41,8 @@ class ProviderPluginRegistrationTests(unittest.TestCase): self.assertEqual(profile.api_mode, "chat_completions") self.assertEqual(profile.base_url, "http://127.0.0.1:8765/v1") self.assertEqual(profile.auth_type, "api_key") - self.assertTrue(profile.supports_vision) + if hasattr(profile, "supports_vision"): + self.assertTrue(profile.supports_vision) self.assertEqual(profile.env_vars, ("CLAUDE_CODE_CLI_API_KEY", "CLAUDE_CODE_CLI_BASE_URL")) self.assertEqual(profile.fallback_models, ("opus", "sonnet", "haiku")) self.assertEqual(profile.default_aux_model, "haiku")