Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions .github/workflows/test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ on:
push:
branches: [main]
pull_request:
workflow_dispatch:

jobs:
unit-tests:
Expand All @@ -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
75 changes: 50 additions & 25 deletions tee_gateway/controllers/chat_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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"
Expand Down
74 changes: 60 additions & 14 deletions tee_gateway/llm_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down
Loading
Loading