feat: 优化用量状态同步机制和 OpenAI CLI 流式转换

- 增加状态回退保护,防止异步响应覆盖已知状态
- 添加轮询并发保护 (pollInFlight),避免重复请求
- 支持 cancelled 状态筛选和显示
- 前端 mergeRecordStatus 保护活跃记录状态
- 后端 streaming 状态同步更新改为使用当前 DB 会话
- OpenAI CLI 流式转换增加工具调用事件支持
- 轮询接口新增 target_model 字段返回
- 修复迁移脚本 inspector 缓存问题,改用 information_schema
This commit is contained in:
fawney19
2026-02-04 03:17:55 +08:00
parent f3e2f84b38
commit 7f09f191a1
17 changed files with 957 additions and 84 deletions

View File

@@ -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":

View File

@@ -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(

View File

@@ -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:
"""根据上下文获取状态字符串"""

View File

@@ -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

View File

@@ -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