mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
feat: 优化用量状态同步机制和 OpenAI CLI 流式转换
- 增加状态回退保护,防止异步响应覆盖已知状态 - 添加轮询并发保护 (pollInFlight),避免重复请求 - 支持 cancelled 状态筛选和显示 - 前端 mergeRecordStatus 保护活跃记录状态 - 后端 streaming 状态同步更新改为使用当前 DB 会话 - OpenAI CLI 流式转换增加工具调用事件支持 - 轮询接口新增 target_model 字段返回 - 修复迁移脚本 inspector 缓存问题,改用 information_schema
This commit is contained in:
@@ -857,7 +857,7 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
|
||||
query = query.filter(Usage.is_stream == False) # noqa: E712
|
||||
elif self.status == "error":
|
||||
query = query.filter((Usage.status_code >= 400) | (Usage.error_message.isnot(None)))
|
||||
elif self.status in ("pending", "streaming", "completed"):
|
||||
elif self.status in ("pending", "streaming", "completed", "cancelled"):
|
||||
# 新的状态筛选:直接按 status 字段过滤
|
||||
query = query.filter(Usage.status == self.status)
|
||||
elif self.status == "failed":
|
||||
|
||||
@@ -2676,7 +2676,28 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
ctx.record_first_byte_time(self.start_time)
|
||||
state["first_yield"] = False
|
||||
if not state["streaming_updated"]:
|
||||
self._update_usage_to_streaming_with_ctx(ctx)
|
||||
# 优先使用当前请求的 DB 会话同步更新,避免状态延迟或丢失
|
||||
try:
|
||||
from src.services.usage import UsageService
|
||||
|
||||
UsageService.update_usage_status(
|
||||
db=self.db,
|
||||
request_id=self.request_id,
|
||||
status="streaming",
|
||||
provider=ctx.provider_name,
|
||||
target_model=ctx.mapped_model,
|
||||
provider_id=ctx.provider_id,
|
||||
provider_endpoint_id=ctx.endpoint_id,
|
||||
provider_api_key_id=ctx.key_id,
|
||||
first_byte_time_ms=ctx.first_byte_time_ms,
|
||||
api_format=ctx.api_format,
|
||||
endpoint_api_format=ctx.provider_api_format or None,
|
||||
has_format_conversion=ctx.needs_conversion,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"[{self.request_id}] 同步更新 streaming 状态失败: {e}")
|
||||
# 回退到后台任务更新
|
||||
self._update_usage_to_streaming_with_ctx(ctx)
|
||||
state["streaming_updated"] = True
|
||||
|
||||
def _convert_sse_line(
|
||||
|
||||
@@ -112,6 +112,7 @@ class StreamTelemetryRecorder:
|
||||
|
||||
try:
|
||||
await self._dispatch_record(
|
||||
bg_db,
|
||||
writer,
|
||||
ctx,
|
||||
original_headers,
|
||||
@@ -137,6 +138,7 @@ class StreamTelemetryRecorder:
|
||||
if response_body is None and should_log_body:
|
||||
response_body = ctx.build_response_body(response_time_ms)
|
||||
await self._dispatch_record(
|
||||
bg_db,
|
||||
db_writer,
|
||||
ctx,
|
||||
original_headers,
|
||||
@@ -438,6 +440,7 @@ class StreamTelemetryRecorder:
|
||||
|
||||
async def _dispatch_record(
|
||||
self,
|
||||
db: Session,
|
||||
writer: TelemetryWriter,
|
||||
ctx: StreamContext,
|
||||
original_headers: dict[str, str],
|
||||
@@ -455,6 +458,14 @@ class StreamTelemetryRecorder:
|
||||
response_body,
|
||||
response_time_ms,
|
||||
)
|
||||
# Queue writer 异步落库可能造成 UI 延迟,先直接更新 Usage 状态
|
||||
if isinstance(writer, QueueTelemetryWriter):
|
||||
await self._update_usage_status_directly(
|
||||
db=db,
|
||||
status=self._get_status_from_ctx(ctx),
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=ctx.status_code,
|
||||
)
|
||||
elif ctx.is_client_disconnected():
|
||||
await self._record_cancelled(
|
||||
writer,
|
||||
@@ -464,6 +475,14 @@ class StreamTelemetryRecorder:
|
||||
response_body,
|
||||
response_time_ms,
|
||||
)
|
||||
# Queue writer 异步落库可能造成 UI 延迟,先直接更新 Usage 状态
|
||||
if isinstance(writer, QueueTelemetryWriter):
|
||||
await self._update_usage_status_directly(
|
||||
db=db,
|
||||
status="cancelled",
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=ctx.status_code,
|
||||
)
|
||||
else:
|
||||
await self._record_failure(
|
||||
writer,
|
||||
@@ -473,6 +492,15 @@ class StreamTelemetryRecorder:
|
||||
response_body,
|
||||
response_time_ms,
|
||||
)
|
||||
# Queue writer 异步落库可能造成 UI 延迟,先直接更新 Usage 状态
|
||||
if isinstance(writer, QueueTelemetryWriter):
|
||||
await self._update_usage_status_directly(
|
||||
db=db,
|
||||
status=self._get_status_from_ctx(ctx),
|
||||
response_time_ms=response_time_ms,
|
||||
status_code=ctx.status_code,
|
||||
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
|
||||
)
|
||||
|
||||
def _get_status_from_ctx(self, ctx: StreamContext) -> str:
|
||||
"""根据上下文获取状态字符串"""
|
||||
|
||||
@@ -411,25 +411,144 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if not state.model:
|
||||
state.model = event.model or ""
|
||||
ss.setdefault("collected_text", "")
|
||||
ss.setdefault("tool_calls", {})
|
||||
ss.setdefault("tool_blocks", {})
|
||||
ss.setdefault("tool_output_index", {})
|
||||
ss.setdefault("output_order", [])
|
||||
ss.setdefault("message_output_index", None)
|
||||
ss.setdefault("message_output_started", False)
|
||||
ss.setdefault("text_started", False)
|
||||
ss.setdefault("next_output_index", 0)
|
||||
ss.setdefault("sent_in_progress", False)
|
||||
response_obj = {
|
||||
"id": state.message_id,
|
||||
"object": "response",
|
||||
"created": int(time.time()),
|
||||
"model": state.model,
|
||||
"status": "in_progress",
|
||||
"output": [],
|
||||
}
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.created",
|
||||
"response": {
|
||||
"id": state.message_id,
|
||||
"object": "response",
|
||||
"created": int(time.time()),
|
||||
"model": state.model,
|
||||
"status": "in_progress",
|
||||
"output": [],
|
||||
},
|
||||
"response": response_obj,
|
||||
}
|
||||
)
|
||||
)
|
||||
# OpenAI Responses API 常见的 in_progress 事件(可选,最佳努力)
|
||||
if not ss.get("sent_in_progress"):
|
||||
ss["sent_in_progress"] = True
|
||||
out.append(event_block({"type": "response.in_progress", "response": response_obj}))
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentBlockStartEvent):
|
||||
# 工具调用块:输出 function_call 添加事件
|
||||
if event.block_type == ContentType.TOOL_USE:
|
||||
tool_id = event.tool_id or ""
|
||||
tool_name = event.tool_name or ""
|
||||
output_index = int(ss.get("next_output_index") or 0)
|
||||
ss["next_output_index"] = output_index + 1
|
||||
if tool_id:
|
||||
tool_calls = ss.setdefault("tool_calls", {})
|
||||
tool_calls.setdefault(tool_id, {"name": tool_name, "args": ""})
|
||||
output_order = ss.setdefault("output_order", [])
|
||||
output_order.append(
|
||||
{"kind": "tool", "id": tool_id, "output_index": output_index}
|
||||
)
|
||||
ss.setdefault("tool_blocks", {})[event.block_index] = tool_id
|
||||
ss.setdefault("tool_output_index", {})[tool_id] = output_index
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"output_index": output_index,
|
||||
"item": {
|
||||
"type": "function_call",
|
||||
"call_id": tool_id,
|
||||
"id": tool_id or f"call_{output_index}",
|
||||
"name": tool_name,
|
||||
"status": "in_progress",
|
||||
"arguments": "",
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, ToolCallDeltaEvent):
|
||||
tool_id = event.tool_id or ss.get("tool_blocks", {}).get(event.block_index, "")
|
||||
if tool_id:
|
||||
tool_calls = ss.setdefault("tool_calls", {})
|
||||
entry = tool_calls.setdefault(tool_id, {"name": "", "args": ""})
|
||||
entry["args"] = str(entry.get("args") or "") + (event.input_delta or "")
|
||||
output_index = ss.get("tool_output_index", {}).get(tool_id, event.block_index)
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.function_call_arguments.delta",
|
||||
"delta": event.input_delta,
|
||||
"item_id": tool_id,
|
||||
"output_index": output_index,
|
||||
}
|
||||
)
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentBlockStopEvent):
|
||||
tool_blocks = ss.get("tool_blocks", {})
|
||||
tool_id = (
|
||||
tool_blocks.pop(event.block_index, None) if isinstance(tool_blocks, dict) else None
|
||||
)
|
||||
if tool_id:
|
||||
tool_calls = ss.get("tool_calls", {})
|
||||
entry = tool_calls.get(tool_id, {})
|
||||
output_index = ss.get("tool_output_index", {}).get(tool_id, event.block_index)
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"output_index": output_index,
|
||||
"item": {
|
||||
"type": "function_call",
|
||||
"call_id": tool_id,
|
||||
"id": tool_id,
|
||||
"name": entry.get("name") or "",
|
||||
"arguments": entry.get("args") or "",
|
||||
"status": "completed",
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
return out
|
||||
|
||||
if isinstance(event, ContentDeltaEvent):
|
||||
if event.text_delta:
|
||||
if not ss.get("message_output_started"):
|
||||
output_index = int(ss.get("next_output_index") or 0)
|
||||
ss["next_output_index"] = output_index + 1
|
||||
ss["message_output_index"] = output_index
|
||||
ss["message_output_started"] = True
|
||||
message_id = f"msg_{state.message_id or 'stream'}"
|
||||
ss.setdefault("output_order", []).append(
|
||||
{"kind": "message", "id": message_id, "output_index": output_index}
|
||||
)
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"output_index": output_index,
|
||||
"item": {
|
||||
"type": "message",
|
||||
"id": message_id,
|
||||
"role": "assistant",
|
||||
"status": "in_progress",
|
||||
"content": [],
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
ss["text_started"] = True
|
||||
ss["collected_text"] = str(ss.get("collected_text") or "") + event.text_delta
|
||||
out.append(
|
||||
event_block(
|
||||
@@ -443,6 +562,34 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
|
||||
if isinstance(event, MessageStopEvent):
|
||||
final_text = str(ss.get("collected_text") or "")
|
||||
message_id = f"msg_{state.message_id or 'stream'}"
|
||||
message_item = {
|
||||
"type": "message",
|
||||
"id": message_id,
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": ([{"type": "output_text", "text": final_text}] if final_text else []),
|
||||
}
|
||||
if ss.get("text_started"):
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.output_text.done",
|
||||
"text": final_text,
|
||||
}
|
||||
)
|
||||
)
|
||||
if ss.get("message_output_started"):
|
||||
output_index = ss.get("message_output_index") or 0
|
||||
out.append(
|
||||
event_block(
|
||||
{
|
||||
"type": "response.output_item.done",
|
||||
"output_index": output_index,
|
||||
"item": message_item,
|
||||
}
|
||||
)
|
||||
)
|
||||
response_obj = self.response_from_internal(
|
||||
InternalResponse(
|
||||
id=state.message_id or "resp",
|
||||
@@ -452,6 +599,53 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
usage=event.usage or UsageInfo(),
|
||||
)
|
||||
)
|
||||
# 将工具调用添加到 output(最佳努力)
|
||||
tool_calls = ss.get("tool_calls", {})
|
||||
output_order = ss.get("output_order", [])
|
||||
output_items: list[dict[str, Any]] = []
|
||||
used_tool_ids: set[str] = set()
|
||||
if isinstance(output_order, list) and output_order:
|
||||
for entry in output_order:
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
if entry.get("kind") == "message":
|
||||
if message_item.get("content"):
|
||||
output_items.append(message_item)
|
||||
elif entry.get("kind") == "tool":
|
||||
tool_id = entry.get("id")
|
||||
if not isinstance(tool_id, str) or not tool_id:
|
||||
continue
|
||||
used_tool_ids.add(tool_id)
|
||||
tool_entry = (
|
||||
tool_calls.get(tool_id) if isinstance(tool_calls, dict) else None
|
||||
)
|
||||
if isinstance(tool_entry, dict):
|
||||
output_items.append(
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": tool_id,
|
||||
"id": tool_id,
|
||||
"name": tool_entry.get("name") or "",
|
||||
"arguments": tool_entry.get("args") or "",
|
||||
"status": "completed",
|
||||
}
|
||||
)
|
||||
if isinstance(tool_calls, dict):
|
||||
for tool_id, tool_entry in tool_calls.items():
|
||||
if tool_id in used_tool_ids or not isinstance(tool_entry, dict):
|
||||
continue
|
||||
output_items.append(
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": tool_id,
|
||||
"id": tool_id,
|
||||
"name": tool_entry.get("name") or "",
|
||||
"arguments": tool_entry.get("args") or "",
|
||||
"status": "completed",
|
||||
}
|
||||
)
|
||||
if output_items:
|
||||
response_obj["output"] = output_items
|
||||
out.append(event_block({"type": "response.completed", "response": response_obj}))
|
||||
return out
|
||||
|
||||
|
||||
@@ -2848,6 +2848,14 @@ class UsageService:
|
||||
logger.warning(f"未找到 request_id={request_id} 的使用记录,无法更新状态")
|
||||
return None
|
||||
|
||||
# 避免状态回退:streaming 只能从 pending/streaming 进入
|
||||
if status == "streaming" and usage.status not in ("pending", "streaming"):
|
||||
logger.debug(
|
||||
f"跳过 streaming 状态更新(避免回退): request_id={request_id}, "
|
||||
f"{usage.status} -> {status}"
|
||||
)
|
||||
return usage
|
||||
|
||||
old_status = usage.status
|
||||
usage.status = status
|
||||
if error_message:
|
||||
@@ -3077,6 +3085,8 @@ class UsageService:
|
||||
Usage.api_format,
|
||||
Usage.endpoint_api_format,
|
||||
Usage.has_format_conversion,
|
||||
# 模型映射(streaming 时已可确定)
|
||||
Usage.target_model,
|
||||
)
|
||||
|
||||
# 管理员轮询:可附带 provider 与上游 key 名称(注意:不要在普通用户接口暴露上游 key 信息)
|
||||
@@ -3234,6 +3244,9 @@ class UsageService:
|
||||
item["endpoint_api_format"] = endpoint_api_format
|
||||
if has_format_conversion is not None:
|
||||
item["has_format_conversion"] = bool(has_format_conversion)
|
||||
# 模型映射(streaming 时已可确定)
|
||||
if r.target_model:
|
||||
item["target_model"] = r.target_model
|
||||
if include_admin_fields:
|
||||
item["provider"] = r.provider_name
|
||||
item["api_key_name"] = r.api_key_name
|
||||
|
||||
Reference in New Issue
Block a user