Spaces:
Running
Running
File size: 13,516 Bytes
8a03d2c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 | """
流式响应处理器(简化版)
实现"假流式":收集所有上游数据后,解析并聚合为一个包含完整上下文(思考过程+最终内容)的响应,一次性发送给下游。
"""
import json
import time
from typing import Any, cast
import collections.abc
from src.core.errors import (
VcoreError,
EmptyResponseError,
InternalError,
NotFoundError,
InvalidArgumentError,
RateLimitError,
)
from src.utils.logger import get_logger
from src.utils.token_counter import calculate_usage_metadata
from .parser import parse_upstream_data
# 初始化日志
logger = get_logger(__name__)
class StreamProcessor:
"""
简化的流式响应处理器 (v3 - 假流式 + 思考过程)
职责:
1. 收集上游所有流式数据
2. 使用 parser 解析聚合后的数据
3. 构建 Gemini 格式的 SSE 响应
"""
def __init__(self):
logger.debug("初始化流处理器")
# 状态追踪
self._actual_content_sent = False
self._request_context: dict[str, Any] = {}
def has_actual_content_sent(self) -> bool:
"""检查是否已发送实际文本内容"""
return self._actual_content_sent
def set_request_context(self, downstream_payload: dict[str, Any], upstream_payload: dict[str, Any]):
"""设置请求上下文"""
logger.debug("设置流处理器请求上下文")
self._request_context = {
'downstream_payload': downstream_payload,
'upstream_payload': upstream_payload
}
def _create_gemini_chunk(
self,
parts: list[dict[str, Any]],
finish_reason: str | None,
safety_ratings: list[dict[str, Any]],
citation_metadata: dict[str, Any],
grounding_metadata: dict[str, Any],
candidate_index: int,
prompt_feedback: dict[str, Any],
usage_metadata: dict[str, Any],
finish_message: str | None = None,
token_count: int | None = None,
avg_logprobs: float | None = None,
logprobs_result: dict[str, Any] | None = None,
create_time: str | None = None,
model_version: str | None = None,
response_id: str | None = None,
model_status: dict[str, Any] | None = None,
) -> str:
"""根据聚合后的内容,创建一个包含完整上下文的Gemini格式SSE事件。"""
candidate: dict[str, Any] = {
"index": candidate_index
}
if finish_reason and isinstance(finish_reason, str):
candidate["finishReason"] = finish_reason.upper()
if finish_message:
candidate["finishMessage"] = finish_message
if parts:
candidate["content"] = {
"parts": parts,
"role": "model"
}
# 只添加实际存在的字段
if safety_ratings:
candidate["safetyRatings"] = safety_ratings
if citation_metadata:
candidate["citationMetadata"] = citation_metadata
if grounding_metadata:
candidate["groundingMetadata"] = grounding_metadata
if token_count is not None:
candidate["tokenCount"] = token_count
if avg_logprobs is not None:
candidate["avgLogprobs"] = avg_logprobs
if logprobs_result:
candidate["logprobsResult"] = logprobs_result
# 根据 doc.md,确保 candidate 其他字段正确
chunk: dict[str, Any] = {"candidates": [candidate]}
# 按照 doc.md 定义
if prompt_feedback:
chunk["promptFeedback"] = prompt_feedback
if usage_metadata:
chunk["usageMetadata"] = usage_metadata
if create_time:
chunk["createTime"] = create_time
if model_version:
chunk["modelVersion"] = model_version
if response_id:
chunk["responseId"] = response_id
if model_status:
chunk["modelStatus"] = model_status
return "data: " + json.dumps(chunk, ensure_ascii=False, separators=(',', ':')) + "\n\n"
async def process_stream(
self,
response_iterator: collections.abc.AsyncIterator[str],
model: str = "vcore-ai-proxy"
) -> collections.abc.AsyncGenerator[str, None]:
"""
处理流式响应(假流式 v3)。
"""
# 移除重复日志,只在开始时记录一次
# logger.info(f"开始处理流式响应: 模型={model}")
start_time = time.time()
raw_chunks: list[str] = []
try:
chunk_count = 0
logger.debug("开始收集上游数据块")
async for chunk in response_iterator:
chunk_count += 1
raw_chunks.append(chunk)
logger.debug(f"收集完成,共 {chunk_count} 个数据块")
raw_data = '\n'.join(raw_chunks)
# 记录完整的原始上游响应
try:
parsed_data = json.loads(raw_data)
logger.debug_json("完整原始上游响应", parsed_data)
except json.JSONDecodeError:
logger.debug_large("完整原始上游响应", raw_data)
if not raw_data:
logger.error("上游返回空数据")
# 触发错误备份
from src.utils.error_logger import save_error_snapshot
save_error_snapshot(
downstream_payload=self._request_context.get('downstream_payload', {}),
upstream_payload=self._request_context.get('upstream_payload', {}),
upstream_response="[EMPTY RESPONSE]",
error_type="empty_response"
)
raise EmptyResponseError("Upstream returned no data")
# 使用独立解析函数
result = parse_upstream_data(raw_data)
# 关键修复:在这里处理解析出的上游错误
if result["has_error"] and not result["parts"]:
error_msg = result["error_message"]
error_obj = result.get("error_obj")
# 排除特定的、不需要备份的错误
is_auth_error = "Failed to verify action" in error_msg or "The caller does not have permission" in error_msg
is_rate_limit = isinstance(error_obj, RateLimitError) or "resource has been exhausted" in error_msg.lower() or "quota" in error_msg.lower()
if not is_auth_error and not is_rate_limit:
logger.error(f"API 错误且无内容: {error_msg}")
# 确定错误类型用于备份目录名
error_type = "api_error"
if error_obj:
# 如果有解析出的错误对象,使用其类名或 code
error_type = f"upstream_{error_obj.code}_{type(error_obj).__name__}"
# 触发错误备份
from src.utils.error_logger import save_error_snapshot
save_error_snapshot(
downstream_payload=self._request_context.get('downstream_payload', {}),
upstream_payload=self._request_context.get('upstream_payload', {}),
upstream_response=raw_data,
error_type=error_type
)
# 如果 parser 已经解析出了错误对象,直接抛出
if error_obj:
raise error_obj
# 降级处理:基于上游错误消息抛出错误
error_msg_lower = error_msg.lower()
if "not found" in error_msg_lower:
raise NotFoundError(message=error_msg)
elif is_rate_limit:
raise RateLimitError(message=error_msg)
elif is_auth_error:
from src.core.errors import AuthenticationError
raise AuthenticationError(
message=f"Authentication/Recaptcha failed: {error_msg}",
details={"upstream_response": error_msg},
upstream_response=error_msg
)
else:
raise InvalidArgumentError(message=error_msg)
finish_reason = result.get("finish_reason") or "STOP"
if not result["parts"] and finish_reason == "STOP" and not result["has_error"]:
if not result.get("prompt_feedback"):
logger.error("上游返回空响应 (无 parts 且 finish_reason=STOP)")
# 触发错误备份
from src.utils.error_logger import save_error_snapshot
save_error_snapshot(
downstream_payload=self._request_context.get('downstream_payload', {}),
upstream_payload=self._request_context.get('upstream_payload', {}),
upstream_response=raw_data,
error_type="stop_no_content"
)
# 快照将由 VcoreAIClient 统一保存
raise EmptyResponseError("Upstream returned empty response (STOP with no content/metadata)")
# 计算 usage metadata
usage_metadata = result.get("usage_metadata", {})
if not usage_metadata and self._request_context:
try:
downstream_payload = self._request_context.get('downstream_payload', {})
# 从请求上下文中提取输入内容
prompt_contents: list[dict[str, Any]] = []
if 'gemini_payload' in downstream_payload:
gemini_payload = downstream_payload['gemini_payload']
if isinstance(gemini_payload, dict) and 'contents' in gemini_payload:
prompt_contents = cast(list[dict[str, Any]], gemini_payload['contents'])
# 计算 token 使用情况
usage_metadata = await calculate_usage_metadata(
prompt_contents=prompt_contents,
response_parts=result["parts"],
request_context=self._request_context
)
except Exception as e:
logger.warning(f"计算 usage metadata 失败: {e}")
usage_metadata = {}
final_chunk = self._create_gemini_chunk(
parts=result["parts"],
finish_reason=result.get("finish_reason"),
safety_ratings=result.get("safety_ratings", []),
citation_metadata=result.get("citation_metadata", {}),
grounding_metadata=result.get("grounding_metadata", {}),
candidate_index=result.get("candidate_index", 0),
prompt_feedback=result.get("prompt_feedback", {}),
usage_metadata=usage_metadata,
finish_message=result.get("finish_message"),
token_count=result.get("token_count"),
avg_logprobs=result.get("avg_logprobs"),
logprobs_result=result.get("logprobs_result"),
create_time=result.get("create_time"),
model_version=result.get("model_version"),
response_id=result.get("response_id"),
model_status=result.get("model_status")
)
process_time = time.time() - start_time
logger.success(f"流式响应处理完成: 耗时={process_time:.3f}s, 完成原因={result.get('finish_reason', 'UNKNOWN')}")
yield final_chunk
self._actual_content_sent = True
except VcoreError as e:
if "Failed to verify action" in e.message or "The caller does not have permission" in e.message:
from src.core.errors import AuthenticationError
if not isinstance(e, AuthenticationError):
raise AuthenticationError(
message=f"Stream contained Authentication error: {e.message}",
details={"upstream_response": e.upstream_response or e.message},
upstream_response=e.upstream_response or e.message
)
else:
logger.error(f"流处理 Vcore 错误: {e.message}")
raise
except Exception as e:
logger.error(f"流处理未知错误: {e}")
# 触发内部错误备份
from src.utils.error_logger import save_error_snapshot
save_error_snapshot(
downstream_payload=self._request_context.get('downstream_payload', {}),
upstream_payload=self._request_context.get('upstream_payload', {}),
upstream_response=str(e),
error_type="internal_exception"
)
raise InternalError(message=f"Unknown stream processing error: {e}")
def get_stream_processor() -> StreamProcessor:
"""创建流处理器实例"""
return StreamProcessor()
|