mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
add missing messages endpoint
This commit is contained in:
+238
-4
@@ -1098,6 +1098,174 @@ 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 httpx.ReadError:
|
||||||
|
if not usage_finalized:
|
||||||
|
await finalize_without_usage()
|
||||||
|
# Upstream dropped the connection mid-stream; response already started, swallow silently
|
||||||
|
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(
|
async def forward_request(
|
||||||
self,
|
self,
|
||||||
request: Request,
|
request: Request,
|
||||||
@@ -1197,7 +1365,54 @@ class BaseUpstreamProvider:
|
|||||||
await client.aclose()
|
await client.aclose()
|
||||||
return mapped_error
|
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"):
|
if path.endswith("chat/completions"):
|
||||||
client_wants_streaming = False
|
client_wants_streaming = False
|
||||||
if request_body:
|
if request_body:
|
||||||
@@ -1825,11 +2040,25 @@ class BaseUpstreamProvider:
|
|||||||
if line.startswith("data: "):
|
if line.startswith("data: "):
|
||||||
try:
|
try:
|
||||||
data_json = json.loads(line[6:])
|
data_json = json.loads(line[6:])
|
||||||
|
# OpenAI format: usage and model at top level
|
||||||
if "usage" in data_json:
|
if "usage" in data_json:
|
||||||
usage_data = data_json["usage"]
|
usage_data = data_json["usage"]
|
||||||
model = data_json.get("model")
|
model = data_json.get("model") or model
|
||||||
elif "model" in data_json and not model:
|
elif "model" in data_json and not model:
|
||||||
model = data_json["model"]
|
model = data_json["model"]
|
||||||
|
# Anthropic format: model and input usage inside "message" key
|
||||||
|
if "message" in data_json:
|
||||||
|
msg = data_json["message"]
|
||||||
|
if not model and msg.get("model"):
|
||||||
|
model = msg["model"]
|
||||||
|
if msg.get("usage") and not usage_data:
|
||||||
|
usage_data = msg["usage"]
|
||||||
|
elif msg.get("usage") and usage_data:
|
||||||
|
# Merge: message_start has input_tokens, message_delta has output_tokens
|
||||||
|
merged = dict(usage_data)
|
||||||
|
for k, v in msg["usage"].items():
|
||||||
|
merged[k] = merged.get(k, 0) + v
|
||||||
|
usage_data = merged
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -2262,9 +2491,14 @@ class BaseUpstreamProvider:
|
|||||||
error_response.headers["X-Cashu"] = refund_token
|
error_response.headers["X-Cashu"] = refund_token
|
||||||
return error_response
|
return error_response
|
||||||
|
|
||||||
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")
|
||||||
|
):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Processing completion/embeddings response",
|
"Processing completion/embeddings/messages response",
|
||||||
extra={"path": path, "amount": amount, "unit": unit},
|
extra={"path": path, "amount": amount, "unit": unit},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user