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