diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 0503808..1cb925f 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -7,6 +7,7 @@ on: push: branches: [main] pull_request: + workflow_dispatch: jobs: unit-tests: @@ -24,3 +25,37 @@ jobs: # To also run integration tests (real CoinGecko network calls), add: # env: # RUN_INTEGRATION_TESTS: "1" + + live-provider-usage: + if: github.event_name == 'workflow_dispatch' + runs-on: ubuntu-latest + timeout-minutes: 20 + env: + RUN_PROVIDER_INTEGRATION_TESTS: "1" + OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }} + ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} + GOOGLE_API_KEY: ${{ secrets.GOOGLE_API_KEY }} + XAI_API_KEY: ${{ secrets.XAI_API_KEY }} + BYTEDANCE_API_KEY: ${{ secrets.BYTEDANCE_API_KEY }} + NOUS_API_KEY: ${{ secrets.NOUS_API_KEY }} + ZAI_API_KEY: ${{ secrets.ZAI_API_KEY }} + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v5 + - name: Check provider API keys + shell: bash + run: | + for name in \ + OPENAI_API_KEY ANTHROPIC_API_KEY GOOGLE_API_KEY XAI_API_KEY \ + BYTEDANCE_API_KEY NOUS_API_KEY ZAI_API_KEY + do + if [[ -z "${!name}" ]]; then + echo "Missing required repository secret: $name" >&2 + exit 1 + fi + done + - name: Run live provider usage tests + run: >- + uv run --group test pytest + tee_gateway/test/test_provider_usage_integration.py + -v --import-mode=importlib diff --git a/tee_gateway/controllers/chat_controller.py b/tee_gateway/controllers/chat_controller.py index 5d042ee..c3249fa 100644 --- a/tee_gateway/controllers/chat_controller.py +++ b/tee_gateway/controllers/chat_controller.py @@ -296,7 +296,11 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): if rf_dict and provider == "anthropic": response = _invoke_anthropic_structured(model, rf_dict, langchain_messages) else: - response = model.invoke(langchain_messages) + # ChatXAI is cached with streaming enabled for the streaming + # endpoint. Disable it here so stream=false gets one complete + # provider response with its authoritative usage object. + invoke_kwargs = {"stream": False} if provider == "x-ai" else {} + response = model.invoke(langchain_messages, **invoke_kwargs) # Normalize content (Gemini may return a list of content parts, and # image-generation models return image blocks alongside any text). @@ -363,7 +367,6 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): f"Response Final\n\tTEE Signature: {signature}\n\tTEE request hash: {input_hash_hex}\n\tTEE output hash: {output_hash_hex}\n\tTEE timestamp: {timestamp}\n\tTEE ID: 0x{tee_keys.get_tee_id()}" ) - # TODO: If no usage is returned, we should compute it here. usage = extract_usage(response) if usage: # Surface the standard OpenAI usage triple on the response; the @@ -381,6 +384,12 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): ) if cost is not None: openai_response["opengradient"] = cost.model_dump(mode="json") + else: + logger.error( + "Provider usage missing model=%s provider=%s source=non_stream", + chat_request.model, + provider, + ) # Validate schema (the extra tee_* fields are preserved by returning dict directly) CreateChatCompletionResponse.from_dict(openai_response) @@ -577,7 +586,19 @@ def generate(): yield f"data: {json.dumps(data)}\n\n" chunks_iter = [] else: - chunks_iter = model.stream(langchain_messages) # type: ignore[assignment] + # xAI Chat Completions requires include_usage for the + # terminal billing event. Its Responses API (used for web + # search) includes terminal usage automatically and does + # not accept the LangChain-only stream_usage argument. + stream_kwargs = ( + {"stream_usage": True} + if provider == "x-ai" and not chat_request.web_search + else {} + ) + chunks_iter = model.stream( # type: ignore[assignment] + langchain_messages, + **stream_kwargs, + ) for chunk in chunks_iter: # Accumulate for post-stream web-search billing (cheap: merges @@ -701,27 +722,26 @@ def generate(): yield f"data: {json.dumps(data)}\n\n" # --- Usage metadata --- - # Accumulate deltas rather than replacing: Gemini returns cumulative - # usageMetadata on every chunk and LangChain emits deltas via - # subtract_usage(), so input_tokens only appears non-zero in the - # *first* chunk carrying usage. Replacing on each chunk would - # overwrite that value with 0 from all subsequent chunks. - if hasattr(chunk, "usage_metadata") and chunk.usage_metadata: - if final_usage is None: - final_usage = {} - for k, v in chunk.usage_metadata.items(): - if isinstance(v, (int, float)): - final_usage[k] = final_usage.get(k, 0) + v - # Thinking tokens (billed at the cheaper text rate) are - # nested in output_token_details and emitted as deltas like - # the top-level counts, so accumulate them the same way. - _otd = chunk.usage_metadata.get("output_token_details") - if isinstance(_otd, dict) and isinstance( - _otd.get("reasoning"), (int, float) - ): - final_usage["reasoning"] = ( - final_usage.get("reasoning", 0) + _otd["reasoning"] - ) + # LangChain emits usage deltas for most providers, so those + # are accumulated. xAI exposes cumulative snapshots instead, + # which are handled separately below. + chunk_usage = extract_usage(chunk) + if chunk_usage: + normalized_chunk_usage = { + "input_tokens": chunk_usage["prompt_tokens"], + "output_tokens": chunk_usage["completion_tokens"], + "total_tokens": chunk_usage["total_tokens"], + "reasoning": chunk_usage.get("reasoning_tokens", 0), + } + if provider == "x-ai": + # xAI sends cumulative usage snapshots. Keep the + # latest one instead of summing every chunk. + final_usage = normalized_chunk_usage + else: + if final_usage is None: + final_usage = {} + for key, value in normalized_chunk_usage.items(): + final_usage[key] = final_usage.get(key, 0) + value # Flush buffered tool calls for OpenAI/Anthropic if buffer_tool_calls and buffered_tool_calls: @@ -786,7 +806,6 @@ def generate(): f"Response Final\n\tTEE Signature: {tee_signature}\n\tTEE request hash: {input_hash_hex}\n\tTEE output hash: {output_hash_hex}\n\tTEE timestamp: {timestamp}\n\tTEE ID: 0x{tee_keys.get_tee_id()}" ) - # TODO: If no usage is returned, we should compute it here. if final_usage: final_data["usage"] = { "prompt_tokens": final_usage.get("input_tokens", 0), @@ -819,6 +838,12 @@ def generate(): f"finish: {finish_reason}, " f"inputHash: {input_hash_hex[:16]}..., outputHash: {output_hash_hex[:16]}..." ) + else: + logger.error( + "Provider usage missing model=%s provider=%s source=stream", + chat_request.model, + provider, + ) yield f"data: {json.dumps(final_data)}\n\n" yield "data: [DONE]\n\n" diff --git a/tee_gateway/llm_backend.py b/tee_gateway/llm_backend.py index b358957..ad31e13 100644 --- a/tee_gateway/llm_backend.py +++ b/tee_gateway/llm_backend.py @@ -8,6 +8,7 @@ import json import logging +from collections.abc import Mapping from typing import List, Dict, Optional, Any, Generator from functools import lru_cache @@ -665,20 +666,65 @@ def convert_messages(messages: list) -> List[Any]: def extract_usage(response) -> Optional[Dict[str, int]]: - """Extract token usage from a LangChain response object.""" - if hasattr(response, "usage_metadata") and response.usage_metadata: - meta = response.usage_metadata - # Thinking tokens, when present, are folded into output_tokens but also - # broken out here. Image-output models bill them at the cheaper - # text/thinking rate (see compute_session_cost), so surface them. - details = meta.get("output_token_details") or {} - return { - "prompt_tokens": meta.get("input_tokens", 0), - "completion_tokens": meta.get("output_tokens", 0), - "total_tokens": meta.get("total_tokens", 0), - "reasoning_tokens": details.get("reasoning", 0), - } - return None + """Extract normalized usage from LangChain or raw provider metadata. + + LangChain normally exposes ``usage_metadata`` using input/output token + names. OpenAI-compatible providers such as xAI can instead leave the raw + ``token_usage`` object in ``response_metadata``. Supporting both shapes + prevents a successful inference from losing its billing data. + """ + + def as_mapping(value: Any) -> Mapping[str, Any] | None: + if isinstance(value, Mapping): + return value + model_dump = getattr(value, "model_dump", None) + if callable(model_dump): + dumped = model_dump() + return dumped if isinstance(dumped, Mapping) else None + return None + + metadata = as_mapping(getattr(response, "usage_metadata", None)) + if metadata is None: + response_metadata = as_mapping(getattr(response, "response_metadata", None)) + if response_metadata is not None: + metadata = as_mapping( + response_metadata.get("token_usage") or response_metadata.get("usage") + ) + + if metadata is None: + return None + + input_tokens = metadata.get("input_tokens", metadata.get("prompt_tokens")) + output_tokens = metadata.get("output_tokens", metadata.get("completion_tokens")) + total_tokens = metadata.get("total_tokens") + if input_tokens is None and output_tokens is None and total_tokens is None: + return None + + prompt_tokens = int(input_tokens or 0) + completion_tokens = int(output_tokens or 0) + total = int( + total_tokens if total_tokens is not None else prompt_tokens + completion_tokens + ) + + reasoning_tokens = 0 + for details_key in ("output_token_details", "output_tokens_details"): + details = as_mapping(metadata.get(details_key)) + if details is not None: + reasoning_tokens = int( + details.get("reasoning", details.get("reasoning_tokens", 0)) or 0 + ) + break + if reasoning_tokens == 0: + details = as_mapping(metadata.get("completion_tokens_details")) + if details is not None: + reasoning_tokens = int(details.get("reasoning_tokens", 0) or 0) + + return { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": total, + "reasoning_tokens": reasoning_tokens, + } # Anthropic's server-side web search tool. The dated type string is the current diff --git a/tee_gateway/test/test_provider_usage_integration.py b/tee_gateway/test/test_provider_usage_integration.py new file mode 100644 index 0000000..b593a8c --- /dev/null +++ b/tee_gateway/test/test_provider_usage_integration.py @@ -0,0 +1,199 @@ +"""Live provider usage and billing integration tests. + +These tests make real, billable requests to one inexpensive chat model from +each supported provider. They are excluded from normal test runs unless +explicitly enabled: + + RUN_PROVIDER_INTEGRATION_TESTS=1 OPENAI_API_KEY=... ... \ + uv run --group test pytest \ + tee_gateway/test/test_provider_usage_integration.py -v + +The tests exercise the same chat-controller and OHTTP billing projection paths +used by the gateway. They do not exercise HPKE or x402 settlement. +""" + +from __future__ import annotations + +import json +import os +import unittest +from dataclasses import dataclass +from decimal import Decimal +from typing import Any, cast +from unittest.mock import MagicMock, patch + +from tee_gateway.config import ProviderConfig +from tee_gateway.controllers import chat_controller, ohttp_controller +from tee_gateway.llm_backend import set_provider_config +from tee_gateway.models import ChatCompletionRequestUserMessage +from tee_gateway.models.create_chat_completion_request import ( + CreateChatCompletionRequest, +) +from tee_gateway.price_feed import set_price_feed + +if os.getenv("RUN_PROVIDER_INTEGRATION_TESTS") != "1": + raise unittest.SkipTest( + "Set RUN_PROVIDER_INTEGRATION_TESTS=1 to run live provider tests" + ) + + +@dataclass(frozen=True) +class _ProviderCase: + provider: str + model: str + secret_name: str + + +PROVIDERS = ( + _ProviderCase("OpenAI", "gpt-4.1-nano", "OPENAI_API_KEY"), + _ProviderCase("Anthropic", "claude-haiku-4-5", "ANTHROPIC_API_KEY"), + _ProviderCase("Google", "gemini-3.5-flash-lite", "GOOGLE_API_KEY"), + _ProviderCase("xAI", "grok-4-fast", "XAI_API_KEY"), + _ProviderCase("ByteDance", "deepseek-v4-flash", "BYTEDANCE_API_KEY"), + _ProviderCase("Nous", "hermes-4-70b", "NOUS_API_KEY"), + _ProviderCase("Z.ai", "glm-5.2", "ZAI_API_KEY"), +) + +_missing_secrets = [ + case.secret_name for case in PROVIDERS if not os.getenv(case.secret_name) +] +if _missing_secrets: + raise RuntimeError( + "Missing provider API keys: " + ", ".join(sorted(_missing_secrets)) + ) + + +class _FixedPriceFeed: + """Deterministic OPG/USD price for testing cost projection.""" + + def get_price(self) -> Decimal: + return Decimal("0.20") + + +def _request(*, model: str, stream: bool) -> CreateChatCompletionRequest: + return CreateChatCompletionRequest( + model=model, + messages=[ + ChatCompletionRequestUserMessage( + role="user", + content="Reply with exactly: OK", + ) + ], + max_tokens=16, + temperature=0, + stream=stream, + ) + + +def _assert_cost_block(test: unittest.TestCase, response: dict) -> None: + usage = cast(dict[str, Any], response.get("usage")) + test.assertIsInstance(usage, dict, response) + test.assertGreater(usage["prompt_tokens"], 0) + test.assertGreater(usage["completion_tokens"], 0) + test.assertGreaterEqual( + usage["total_tokens"], + usage["prompt_tokens"] + usage["completion_tokens"], + ) + + cost = cast(dict[str, Any], response.get("opengradient")) + test.assertIsInstance(cost, dict, response) + for field in ("cost_opg", "cost_usd", "opg_price_usd"): + test.assertIn(field, cost, response) + test.assertNotIn(cost[field], (None, ""), response) + test.assertGreater(int(cost["cost_opg"]), 0) + test.assertGreater(Decimal(cost["cost_usd"]), 0) + test.assertGreater(Decimal(cost["opg_price_usd"]), 0) + + +def _stream_events(response) -> list[dict]: + events: list[dict] = [] + done_seen = False + for chunk in response.response: + text = chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk + for line in text.splitlines(): + if not line.startswith("data:"): + continue + payload = line[len("data:") :].strip() + if payload == "[DONE]": + done_seen = True + elif payload: + events.append(json.loads(payload)) + if not done_seen: + raise AssertionError("Streaming response did not emit [DONE]") + return events + + +class TestLiveProviderUsageBilling(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + set_provider_config( + ProviderConfig( + openai_api_key=os.environ["OPENAI_API_KEY"], + anthropic_api_key=os.environ["ANTHROPIC_API_KEY"], + google_api_key=os.environ["GOOGLE_API_KEY"], + xai_api_key=os.environ["XAI_API_KEY"], + bytedance_api_key=os.environ["BYTEDANCE_API_KEY"], + nous_api_key=os.environ["NOUS_API_KEY"], + zai_api_key=os.environ["ZAI_API_KEY"], + ) + ) + set_price_feed(_FixedPriceFeed()) # type: ignore[arg-type] + + def setUp(self) -> None: + self.tee_keys = MagicMock() + self.tee_keys.sign_data.return_value = "integration-test-signature" + self.tee_keys.get_tee_id.return_value = "00" * 32 + + def test_non_streaming_usage_and_cost_for_every_provider(self) -> None: + for case in PROVIDERS: + with self.subTest(provider=case.provider, model=case.model): + with patch.object( + chat_controller, + "get_tee_keys", + return_value=self.tee_keys, + ): + response = chat_controller._create_non_streaming_response( + _request(model=case.model, stream=False) + ) + + self.assertIsInstance(response, dict, response) + self.assertNotIn("error", response) + _assert_cost_block(self, response) + + headers = ohttp_controller._extract_cost_headers( + json.dumps(response).encode("utf-8") + ) + self.assertEqual( + headers["X-Inference-Cost-USD"], + response["opengradient"]["cost_usd"], + ) + self.assertEqual( + headers["X-Inference-Cost-OPG"], + response["opengradient"]["cost_opg"], + ) + + def test_streaming_usage_and_cost_for_every_provider(self) -> None: + for case in PROVIDERS: + with self.subTest(provider=case.provider, model=case.model): + with patch.object( + chat_controller, + "get_tee_keys", + return_value=self.tee_keys, + ): + response = chat_controller._create_streaming_response( + _request(model=case.model, stream=True) + ) + events = _stream_events(response) + + final = events[-1] + self.assertNotIn("error", final) + _assert_cost_block(self, final) + + billing_frame = ohttp_controller._build_billing_frame(final) + self.assertTrue( + billing_frame.startswith(ohttp_controller.OHTTP_BILLING_FRAME_MAGIC) + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tee_gateway/test/test_tee_core.py b/tee_gateway/test/test_tee_core.py index 55f87aa..f772597 100644 --- a/tee_gateway/test/test_tee_core.py +++ b/tee_gateway/test/test_tee_core.py @@ -901,6 +901,60 @@ def test_returns_none_when_attribute_missing(self): mock_resp = type("R", (), {})() self.assertIsNone(extract_usage(mock_resp)) + def test_falls_back_to_raw_chat_completion_usage(self): + mock_resp = type( + "R", + (), + { + "usage_metadata": None, + "response_metadata": { + "token_usage": { + "prompt_tokens": 11, + "completion_tokens": 7, + "total_tokens": 18, + "completion_tokens_details": {"reasoning_tokens": 3}, + } + }, + }, + )() + usage = extract_usage(mock_resp) + self.assertEqual( + usage, + { + "prompt_tokens": 11, + "completion_tokens": 7, + "total_tokens": 18, + "reasoning_tokens": 3, + }, + ) + + def test_falls_back_to_raw_responses_usage(self): + mock_resp = type( + "R", + (), + { + "usage_metadata": None, + "response_metadata": { + "usage": { + "input_tokens": 13, + "output_tokens": 5, + "total_tokens": 18, + "output_tokens_details": {"reasoning_tokens": 2}, + } + }, + }, + )() + usage = extract_usage(mock_resp) + self.assertEqual( + usage, + { + "prompt_tokens": 13, + "completion_tokens": 5, + "total_tokens": 18, + "reasoning_tokens": 2, + }, + ) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_opengradient_field.py b/tests/test_opengradient_field.py index a26c88c..db2bdb7 100644 --- a/tests/test_opengradient_field.py +++ b/tests/test_opengradient_field.py @@ -79,6 +79,42 @@ def test_opengradient_block_embedded_when_cost_computed(self): self.assertIn("opengradient", resp) self.assertEqual(resp["opengradient"]["cost_opg"], "12345") + def test_xai_non_streaming_disables_provider_streaming(self): + from tee_gateway.controllers import chat_controller + + fake_response = MagicMock() + fake_response.content = "hello" + fake_response.tool_calls = None + fake_model = MagicMock() + fake_model.invoke.return_value = fake_response + request = _chat_request() + request.model = "grok-4.3" + + with ( + patch.object( + chat_controller, "get_chat_model_cached", return_value=fake_model + ), + patch.object( + chat_controller, "get_provider_from_model", return_value="x-ai" + ), + patch.object(chat_controller, "extract_usage", return_value=_FAKE_USAGE), + patch.object( + chat_controller, "compute_session_cost", return_value=_fake_cost() + ), + patch.object( + chat_controller, "get_tee_keys", return_value=_fake_tee_keys() + ), + patch.object( + chat_controller, + "compute_tee_msg_hash", + return_value=(b"h", "ih", "oh"), + ), + ): + chat_controller._create_non_streaming_response(request) + + fake_model.invoke.assert_called_once() + self.assertFalse(fake_model.invoke.call_args.kwargs["stream"]) + def test_opengradient_block_absent_when_compute_returns_none(self): from tee_gateway.controllers import chat_controller @@ -157,6 +193,92 @@ def test_final_sse_event_carries_opengradient(self): self.assertIn("opengradient", final) self.assertEqual(final["opengradient"]["cost_opg"], "12345") + def test_xai_stream_explicitly_requests_usage(self): + from tee_gateway.controllers import chat_controller + + chunk = MagicMock() + chunk.content = "hello" + chunk.tool_call_chunks = [] + chunk.usage_metadata = { + "input_tokens": 10, + "output_tokens": 20, + "total_tokens": 30, + } + fake_model = MagicMock() + fake_model.stream.return_value = iter([chunk]) + + with ( + patch.object( + chat_controller, "get_chat_model_cached", return_value=fake_model + ), + patch.object( + chat_controller, "get_provider_from_model", return_value="x-ai" + ), + patch.object( + chat_controller, "compute_session_cost", return_value=_fake_cost() + ), + patch.object( + chat_controller, "get_tee_keys", return_value=_fake_tee_keys() + ), + patch.object( + chat_controller, + "compute_tee_msg_hash", + return_value=(b"h", "ih", "oh"), + ), + ): + request = _chat_request() + request.model = "grok-4.3" + request.stream = True + response = chat_controller._create_streaming_response(request) + list(response.response) + + fake_model.stream.assert_called_once() + self.assertTrue(fake_model.stream.call_args.kwargs["stream_usage"]) + + def test_xai_responses_stream_does_not_forward_stream_usage(self): + from tee_gateway.controllers import chat_controller + + chunk = MagicMock() + chunk.content = "hello" + chunk.tool_call_chunks = [] + chunk.usage_metadata = { + "input_tokens": 10, + "output_tokens": 20, + "total_tokens": 30, + } + fake_model = MagicMock() + fake_model.stream.return_value = iter([chunk]) + + with ( + patch.object( + chat_controller, "get_chat_model_cached", return_value=fake_model + ), + patch.object( + chat_controller, "get_provider_from_model", return_value="x-ai" + ), + patch.object(chat_controller, "_build_tools_list", return_value=[]), + patch.object( + chat_controller, "compute_session_cost", return_value=_fake_cost() + ), + patch.object( + chat_controller, "get_tee_keys", return_value=_fake_tee_keys() + ), + patch.object( + chat_controller, + "compute_tee_msg_hash", + return_value=(b"h", "ih", "oh"), + ), + ): + request = _chat_request() + request.model = "grok-4.3" + request.stream = True + request.web_search = True + response = chat_controller._create_streaming_response(request) + list(response.response) + + fake_model.stream.assert_called_once() + self.assertNotIn("stream_usage", fake_model.stream.call_args.kwargs) + class TestCompletionsOpengradient(unittest.TestCase): def test_opengradient_block_embedded_on_completion(self):