From b9879cbea706345d48d16a43b38ff0fa3e23ea9a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 10 Jan 2026 23:32:45 +0100 Subject: [PATCH] claude code integration --- routstr/upstream/base.py | 213 ++++++++++++++++++++++++++++++++++++++- 1 file changed, 212 insertions(+), 1 deletion(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 1c42da8c..9d461990 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -1064,6 +1064,170 @@ class BaseUpstreamProvider: }, ) + async def handle_streaming_messages_completion( + self, response: httpx.Response, key: ApiKey, max_cost_for_model: int + ) -> StreamingResponse: + async def stream_with_cost( + max_cost_for_model: int, + ) -> AsyncGenerator[bytes, None]: + stored_chunks: list[bytes] = [] + usage_finalized: bool = False + last_model_seen: str | None = None + input_tokens: int = 0 + output_tokens: int = 0 + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + usage_finalized = True + return None + try: + fallback: dict = { + "model": last_model_seen or "unknown", + "usage": None, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, fallback, new_session, max_cost_for_model + ) + usage_finalized = True + return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception: + usage_finalized = True + return None + + try: + async for chunk in response.aiter_bytes(): + stored_chunks.append(chunk) + try: + decoded_chunk = chunk.decode("utf-8", errors="ignore") + for line in decoded_chunk.split("\n"): + if line.startswith("data: "): + try: + data = json.loads(line[6:]) + if isinstance(data, dict): + msg = data.get("message", {}) + if msg and msg.get("model"): + last_model_seen = str(msg.get("model")) + + if usage := msg.get("usage"): + input_tokens += usage.get("input_tokens", 0) + output_tokens += usage.get( + "output_tokens", 0 + ) + + if usage := data.get("usage"): + input_tokens += usage.get("input_tokens", 0) + output_tokens += usage.get( + "output_tokens", 0 + ) + except json.JSONDecodeError: + pass + except Exception: + pass + + yield chunk + + usage_data = { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + } + + if input_tokens > 0 or output_tokens > 0: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if fresh_key: + try: + combined_data = { + "model": last_model_seen or "unknown", + "usage": usage_data, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, + combined_data, + new_session, + max_cost_for_model, + ) + usage_finalized = True + yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception: + pass + + if not usage_finalized: + maybe_cost_event = await finalize_without_usage() + if maybe_cost_event is not None: + yield maybe_cost_event + + except Exception: + if not usage_finalized: + await finalize_without_usage() + raise + finally: + if not usage_finalized: + await finalize_without_usage() + + response_headers = dict(response.headers) + response_headers.pop("content-encoding", None) + response_headers.pop("content-length", None) + + return StreamingResponse( + stream_with_cost(max_cost_for_model), + status_code=response.status_code, + headers=response_headers, + ) + + async def handle_non_streaming_messages_completion( + self, + response: httpx.Response, + key: ApiKey, + session: AsyncSession, + deducted_max_cost: int, + path: str, + ) -> Response: + try: + content = await response.aread() + response_json = json.loads(content) + + if path.endswith("count_tokens") and "usage" not in response_json: + input_tokens = response_json.get("input_tokens", 0) + response_json["usage"] = {"input_tokens": input_tokens} + + cost_data = await adjust_payment_for_tokens( + key, response_json, session, deducted_max_cost + ) + response_json["cost"] = cost_data + + allowed_headers = { + "content-type", + "cache-control", + "date", + "vary", + "access-control-allow-origin", + "access-control-allow-methods", + "access-control-allow-headers", + "access-control-allow-credentials", + "access-control-expose-headers", + "access-control-max-age", + } + + response_headers = { + k: v + for k, v in response.headers.items() + if k.lower() in allowed_headers + } + + return Response( + content=json.dumps(response_json).encode(), + status_code=response.status_code, + headers=response_headers, + media_type="application/json", + ) + except Exception: + raise + async def forward_request( self, request: Request, @@ -1157,7 +1321,54 @@ class BaseUpstreamProvider: await client.aclose() return mapped_error - if path.endswith("chat/completions") or path.endswith("embeddings"): + if ( + path.endswith("chat/completions") + or path.endswith("embeddings") + or path.endswith("messages") + or path.endswith("messages/count_tokens") + ): + if path.endswith("messages"): + client_wants_streaming = False + if request_body: + try: + request_data = json.loads(request_body) + client_wants_streaming = request_data.get("stream", False) + except json.JSONDecodeError: + pass + + content_type = response.headers.get("content-type", "") + upstream_is_streaming = "text/event-stream" in content_type + is_streaming = client_wants_streaming and upstream_is_streaming + + if is_streaming and response.status_code == 200: + result = await self.handle_streaming_messages_completion( + response, key, max_cost_for_model + ) + background_tasks = BackgroundTasks() + background_tasks.add_task(response.aclose) + background_tasks.add_task(client.aclose) + result.background = background_tasks + return result + + if response.status_code == 200: + try: + return await self.handle_non_streaming_messages_completion( + response, key, session, max_cost_for_model, path + ) + finally: + await response.aclose() + await client.aclose() + + if path.endswith("messages/count_tokens"): + if response.status_code == 200: + try: + return await self.handle_non_streaming_messages_completion( + response, key, session, max_cost_for_model, path + ) + finally: + await response.aclose() + await client.aclose() + if path.endswith("chat/completions"): client_wants_streaming = False if request_body: