ai-model-studio / src /stream /parser.py
slk1st's picture
Deploy AI Model Studio
8a03d2c verified
Raw
History Blame Contribute Delete
18.9 kB
"""
Vcore AI 响应解析工具函数
负责解析上游 API 的响应数据,处理错误和元数据提取。
"""
import json
from typing import Any, cast
from src.core.errors import (
VcoreError,
InternalError,
parse_error_response,
)
from src.utils.logger import get_logger
from src.utils.string_utils import snake_to_camel
logger = get_logger(__name__)
def _coerce_function_args(args: Any) -> dict[str, Any]:
if isinstance(args, dict):
return cast(dict[str, Any], args)
if isinstance(args, str):
try:
parsed = json.loads(args)
return parsed if isinstance(parsed, dict) else {"value": parsed}
except json.JSONDecodeError:
return {"raw": args}
if args is None:
return {}
return {"value": args}
def _coerce_function_response(response: Any) -> dict[str, Any]:
if isinstance(response, dict):
return cast(dict[str, Any], response)
if isinstance(response, str):
try:
parsed = json.loads(response)
return parsed if isinstance(parsed, dict) else {"result": parsed}
except json.JSONDecodeError:
return {"result": response}
if response is None:
return {}
return {"result": response}
def extract_path_index(result: dict[str, Any]) -> int:
"""从 result 对象中提取 path 索引"""
path = result.get('path', [])
if not path or not isinstance(path, list):
return -1
try:
# 索引通常是路径的最后一个整数元素
path_list = cast(list[Any], path)
for elem in reversed(path_list):
if isinstance(elem, int):
return elem
if isinstance(elem, str) and elem.isdigit():
return int(elem)
except (IndexError, ValueError, TypeError):
pass
return -1
def clean_json_string(raw_data: str) -> str:
"""清理并规范化 JSON 字符串"""
cleaned_data = raw_data.strip()
if cleaned_data.endswith(','):
cleaned_data = cleaned_data[:-1]
if not cleaned_data.startswith('['):
cleaned_data = f'[{cleaned_data}]'
elif not cleaned_data.endswith(']'):
# 这是一个不完整的数组,尝试补全
if '}]' not in cleaned_data:
cleaned_data += ']'
return cleaned_data
def process_candidate_metadata(candidate_data: dict[str, Any]) -> dict[str, Any]:
"""提取 candidate 级别的元数据(仅提取实际存在的字段)"""
metadata: dict[str, Any] = {}
finish_reason = candidate_data.get('finishReason')
if finish_reason:
metadata['finish_reason'] = finish_reason
if 'finishMessage' in candidate_data:
metadata['finish_message'] = candidate_data['finishMessage']
# 提取实际存在的元数据字段
if candidate_data.get('safetyRatings'):
metadata['safety_ratings'] = candidate_data['safetyRatings']
if candidate_data.get('citationMetadata'):
metadata['citation_metadata'] = candidate_data['citationMetadata']
if candidate_data.get('groundingMetadata'):
metadata['grounding_metadata'] = candidate_data['groundingMetadata']
if 'tokenCount' in candidate_data:
metadata['token_count'] = candidate_data['tokenCount']
if 'avgLogprobs' in candidate_data:
metadata['avg_logprobs'] = candidate_data['avgLogprobs']
if 'logprobsResult' in candidate_data:
metadata['logprobs_result'] = candidate_data['logprobsResult']
# index 字段通常存在
if candidate_data.get('index') is not None:
metadata['candidate_index'] = candidate_data['index']
return metadata
def _extract_error_message(item: dict[str, Any]) -> str | None:
"""从单个响应项中提取错误信息(如果有)"""
# 1. 检查顶层 error 对象 (标准 Google Cloud 错误)
error_obj = item.get('error')
if error_obj:
if isinstance(error_obj, dict):
# 显式转换为 Dict 以满足类型检查
safe_error_obj = cast(dict[str, Any], error_obj)
return str(safe_error_obj.get('message', str(safe_error_obj)))
return str(error_obj)
# 2. 检查顶层 errors 列表 (GraphQL 风格或批处理错误)
errors = item.get('errors')
if errors and isinstance(errors, list):
# 显式转换
safe_errors = cast(list[Any], errors)
if safe_errors:
first_error = safe_errors[0]
if isinstance(first_error, dict):
safe_first_error = cast(dict[str, Any], first_error)
return str(safe_first_error.get('message', str(safe_first_error)))
return str(first_error)
return None
def _clean_part_fields(part: dict[str, Any]) -> dict[str, Any]:
"""
清理 part 中的空字段,只保留有实际内容的字段
Args:
part: 原始 part 字典
Returns:
清理后的 part 字典
"""
if 'function_call' in part and 'functionCall' not in part:
part = {**part, 'functionCall': part['function_call']}
if 'function_response' in part and 'functionResponse' not in part:
part = {**part, 'functionResponse': part['function_response']}
if 'thought_signature' in part and 'thoughtSignature' not in part:
part = {**part, 'thoughtSignature': part['thought_signature']}
if 'inline_data' in part and 'inlineData' not in part:
inline_data = part['inline_data']
if isinstance(inline_data, dict):
part = {**part, 'inlineData': {snake_to_camel(str(k)): v for k, v in cast(dict[Any, Any], inline_data).items()}}
else:
part = {**part, 'inlineData': inline_data}
if 'file_data' in part and 'fileData' not in part:
file_data = part['file_data']
if isinstance(file_data, dict):
part = {**part, 'fileData': {snake_to_camel(str(k)): v for k, v in cast(dict[Any, Any], file_data).items()}}
else:
part = {**part, 'fileData': file_data}
try:
from src.core.types import ContentPart
content_part = ContentPart.model_validate(part)
cleaned_part = content_part.model_dump(exclude_none=True, by_alias=True)
# 额外处理 text 为空字符串的情况,以及一些需要保留真实值的嵌套字典
if 'text' in cleaned_part and str(cleaned_part['text']) == "":
del cleaned_part['text']
if 'thought_signature' in cleaned_part and 'thoughtSignature' not in cleaned_part:
cleaned_part['thoughtSignature'] = cleaned_part.pop('thought_signature')
return cleaned_part
except Exception as e:
logger.debug(f"Pydantic 验证 part 失败,回退到基础清洗: {e}")
cleaned_part: dict[str, Any] = {}
# 处理文本内容
if 'text' in part:
text_value = part['text']
if text_value is not None and str(text_value) != "":
cleaned_part['text'] = text_value
# 处理思考标记
if 'thought' in part:
cleaned_part['thought'] = part['thought']
# 处理思考签名,兼容 REST snake_case 和 SDK camelCase
if 'thoughtSignature' in part:
cleaned_part['thoughtSignature'] = part['thoughtSignature']
elif 'thought_signature' in part:
cleaned_part['thoughtSignature'] = part['thought_signature']
# 处理数据类型标记(如果存在且不为空)
if 'data' in part and part['data']:
cleaned_part['data'] = part['data']
# 处理函数调用(只保留有名称的)
if 'functionCall' in part or 'function_call' in part:
func_call = part.get('functionCall') or part.get('function_call')
if isinstance(func_call, dict):
func_call_dict = cast(dict[str, Any], func_call)
if func_call_dict.get('name') and str(func_call_dict['name']).strip():
fixed_call = func_call_dict.copy()
fixed_call['args'] = _coerce_function_args(fixed_call.get('args', {}))
cleaned_part['functionCall'] = fixed_call
# 处理函数响应(只保留有名称的)
if 'functionResponse' in part or 'function_response' in part:
func_response = part.get('functionResponse') or part.get('function_response')
if isinstance(func_response, dict):
func_response_dict = cast(dict[str, Any], func_response)
current_name = func_response_dict.get('name') or func_response_dict.get('functionName') or func_response_dict.get('function_name')
if current_name and str(current_name).strip():
fixed_response = func_response_dict.copy()
fixed_response['name'] = str(current_name)
fixed_response['response'] = _coerce_function_response(fixed_response.get('response', {}))
cleaned_part['functionResponse'] = fixed_response
# 处理内联数据(只保留有实际数据的)
if 'inlineData' in part:
inline_data = part['inlineData']
if isinstance(inline_data, dict):
inline_data_dict = cast(dict[str, Any], inline_data)
if (inline_data_dict.get('data') and
str(inline_data_dict['data']).strip() and
inline_data_dict.get('mimeType') and
str(inline_data_dict['mimeType']).strip()):
cleaned_part['inlineData'] = inline_data
# 处理文件数据(只保留有实际数据的)
if 'fileData' in part:
file_data = part['fileData']
if isinstance(file_data, dict):
file_data_dict = cast(dict[str, Any], file_data)
if (file_data_dict.get('fileUri') and
str(file_data_dict['fileUri']).strip() and
file_data_dict.get('mimeType') and
str(file_data_dict['mimeType']).strip()):
cleaned_part['fileData'] = file_data
return cleaned_part
def _merge_content_blocks(parts: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""
合并思考块和非思考块的文本内容
Args:
parts: 原始的 parts 列表
Returns:
合并后的 parts 列表,思考块在前,非思考块在后
"""
# 首先清理所有 parts 中的空字段
cleaned_parts = [_clean_part_fields(part) for part in parts]
# 过滤掉完全为空的 parts
cleaned_parts = [part for part in cleaned_parts if part]
thought_texts: list[str] = [] # 思考块文本
thought_signatures: list[str] = [] # 思考块签名
content_texts: list[str] = [] # 非思考块文本
other_parts: list[dict[str, Any]] = [] # 非文本部分(函数调用等)
# 分类处理所有 parts
for part in cleaned_parts:
# 检查是否为文本块
if 'text' in part and part['text'] is not None:
text_content = str(part['text'])
if text_content == "":
continue
# 判断是否为思考块
is_thought = part.get('thought', False)
if is_thought:
thought_texts.append(text_content)
if 'thoughtSignature' in part:
thought_signatures.append(part['thoughtSignature'])
else:
content_texts.append(text_content)
else:
# 非文本部分(如函数调用、函数响应等)保持原样
other_parts.append(part)
# 构建合并后的 parts 列表
merged_parts: list[dict[str, Any]] = []
# 1. 合并思考块文本(如果有)
if thought_texts:
merged_thought_text = ''.join(thought_texts)
thought_part = {
'text': merged_thought_text,
'thought': True
}
# 重新插入签名(使用最后一个有效的签名,或者可以根据需要合并)
if thought_signatures:
thought_part['thoughtSignature'] = thought_signatures[-1]
merged_parts.append(thought_part)
# 2. 添加非文本部分(保持原有顺序)
merged_parts.extend(other_parts)
# 3. 合并非思考块文本(如果有)
if content_texts:
merged_content_text = ''.join(content_texts)
merged_parts.append({
'text': merged_content_text
})
return merged_parts
def parse_upstream_data(raw_data: str) -> dict[str, Any]:
"""
解析完整的上游原始数据。
Returns:
包含 parts, finish_reason 和实际存在的元数据的字典
"""
state: dict[str, Any] = {
"finish_reason": None,
"finish_message": None,
"safety_ratings": [],
"citation_metadata": {},
"grounding_metadata": {},
"token_count": None,
"avg_logprobs": None,
"logprobs_result": None,
"candidate_index": 0,
"prompt_feedback": {},
"usage_metadata": {},
"create_time": None,
"model_version": None,
"response_id": None,
"model_status": None,
"has_error": False,
"error_message": "",
"error_obj": None,
"parts_by_path": {},
"unindexed_parts": []
}
try:
cleaned_data = clean_json_string(raw_data)
data_list = json.loads(cleaned_data)
if not isinstance(data_list, list):
data_list = [data_list]
# 显式转换为 List[Any]
safe_data_list = cast(list[Any], data_list)
for item in safe_data_list:
if not isinstance(item, dict):
continue
item_dict = cast(dict[str, Any], item)
# 1. 优先通过统一解析逻辑检查 item 中的错误 (处理 errors 数组)
parsed_error = parse_error_response(item_dict)
if parsed_error:
# 如果是 "Failed to verify action",我们忽略它,因为这可能是匿名的第一个预期错误
if "Failed to verify action" in parsed_error.message:
logger.debug(f"忽略预期的认证错误: {parsed_error.message}")
else:
state["has_error"] = True
state["error_message"] = parsed_error.message
state["error_obj"] = parsed_error
# 继续解析以尝试提取更多上下文
# 2. 检查顶层错误 (作为兜底)
error_msg = _extract_error_message(item_dict)
if error_msg and not state["has_error"]:
state["has_error"] = True
state["error_message"] = error_msg
# 3. 处理 results 列表
results = item_dict.get('results', [])
if not isinstance(results, list):
continue
typed_results: list[dict[str, Any]] = []
safe_results = cast(list[Any], results)
for r in safe_results:
if isinstance(r, dict):
typed_results.append(cast(dict[str, Any], r))
# 3. 检查 results 中的错误 (使用统一解析逻辑)
parsed_error = parse_error_response(typed_results)
if parsed_error:
state["has_error"] = True
state["error_message"] = parsed_error.message
state["error_obj"] = parsed_error
# 4. 提取数据 parts
for result in typed_results:
# 如果有 data=null 且有 errors,这已经被上面的 parsed_error 捕获
# 我们跳过这个 result 的 data 处理,避免 NoneType 错误
if result.get('data') is None and 'errors' in result:
continue
path_index = extract_path_index(result)
data = result.get('data')
if isinstance(data, dict):
_update_state_from_data(state, cast(dict[str, Any], data), path_index)
except json.JSONDecodeError as e:
state["has_error"] = True
state["error_message"] = f"JSON parse error: {e}"
except VcoreError:
raise
except Exception as e:
# 捕获其他未预期的解析错误
logger.error(f"解析过程发生未知错误: {e}")
state["has_error"] = True
state["error_message"] = f"Parse error: {str(e)}"
# 组装 parts - 先按原有逻辑收集所有parts
parts_by_path = cast(dict[int, Any], state['parts_by_path'])
ordered_parts: list[dict[str, Any]] = [parts_by_path[k] for k in sorted(parts_by_path.keys())]
unindexed_parts = cast(list[Any], state['unindexed_parts'])
ordered_parts.extend(unindexed_parts)
# 新增:合并思考块和非思考块
final_parts = _merge_content_blocks(ordered_parts)
result: dict[str, Any] = {
"parts": final_parts
}
# 将 state 中除了临时存储结构外的所有字段合并到 result
excluded_keys = ['parts_by_path', 'unindexed_parts']
result.update({k: v for k, v in state.items() if k not in excluded_keys})
return result
def _update_state_from_data(state: dict[str, Any], data: dict[str, Any], path_index: int):
"""从数据对象更新解析状态(仅提取实际存在的字段)"""
# 只提取实际存在的顶层元数据
if data.get('promptFeedback'):
state['prompt_feedback'] = data['promptFeedback']
if 'usageMetadata' in data:
state['usage_metadata'] = data['usageMetadata']
if 'createTime' in data:
state['create_time'] = data['createTime']
if 'modelVersion' in data:
state['model_version'] = data['modelVersion']
if 'responseId' in data:
state['response_id'] = data['responseId']
if 'modelStatus' in data:
state['model_status'] = data['modelStatus']
# 处理 candidates
candidates = data.get('candidates', [])
for candidate in candidates:
# 提取 candidate 元数据 (包括 finish_reason)
meta = process_candidate_metadata(candidate)
for k, v in meta.items():
if v is not None and v != [] and v != {}:
state[k] = v
# 提取 content parts
content = candidate.get('content', {})
parts = content.get('parts', [])
for part in parts:
if path_index != -1:
state['parts_by_path'][path_index] = part
else:
state['unindexed_parts'].append(part)