mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
refactor: 引入模块钩子系统,解耦认证逻辑,支持模块/normalizer/parser 自动发现
- 新增 HookDispatcher 钩子分发器,支持 FIRST_RESULT 和 COLLECT_ALL 两种策略 - LDAP 认证逻辑从 AuthService 移至 ldap 模块钩子实现 - Management Token 前缀认证从 pipeline 硬编码改为模块钩子注册 - src/modules/ 改为自动扫描子目录发现 ModuleDefinition - normalizers 和 parsers 注册改为基于类属性自动发现 - OpenAI CLI 增加 /v1/responses/compact 端点和并行 tool_call 支持 - OpenAI CLI normalizer 支持 Chat Completions 格式自动回退 - Codex 适配器增加 compact 模式上下文传递和 header 调整 - HeaderBuilder 改进非 latin-1 字符处理(UTF-8 字节透传) - Gunicorn 增加 graceful_timeout 防止僵尸进程
This commit is contained in:
@@ -609,21 +609,33 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
# tool_calls delta
|
||||
tool_calls = delta.get("tool_calls")
|
||||
if isinstance(tool_calls, list):
|
||||
# index -> tool_id 映射(用于后续 delta 缺少 id 时查找)
|
||||
index_to_id: dict[str, str] = ss.setdefault("tool_index_to_id", {})
|
||||
|
||||
for tool_call in tool_calls:
|
||||
if not isinstance(tool_call, dict):
|
||||
continue
|
||||
|
||||
tc_id = str(tool_call.get("id") or "")
|
||||
raw_index = tool_call.get("index")
|
||||
tc_index = str(raw_index) if raw_index is not None else ""
|
||||
fn = tool_call.get("function") or {}
|
||||
fn = fn if isinstance(fn, dict) else {}
|
||||
tc_name = str(fn.get("name") or "")
|
||||
tc_args = fn.get("arguments")
|
||||
|
||||
block_index = self._ensure_tool_block_index(
|
||||
ss, tc_id or str(tool_call.get("index") or "")
|
||||
)
|
||||
# 首次出现的 delta 同时有 id 和 index,记录映射
|
||||
if tc_id and tc_index:
|
||||
index_to_id[tc_index] = tc_id
|
||||
# 后续 delta 只有 index 没有 id,通过映射恢复 id
|
||||
elif not tc_id and tc_index and tc_index in index_to_id:
|
||||
tc_id = index_to_id[tc_index]
|
||||
|
||||
# tool start(只在首次见到该 tool_id 时发)
|
||||
# 用 tool_id 作为优先 key(确保同一 tool call 始终同一 block_index)
|
||||
tool_key = tc_id or tc_index
|
||||
block_index = self._ensure_tool_block_index(ss, tool_key)
|
||||
|
||||
# tool start(只在首次见到该 tool_key 时发)
|
||||
started_key = f"tool_started:{block_index}"
|
||||
if not ss.get(started_key):
|
||||
ss[started_key] = True
|
||||
@@ -649,7 +661,7 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
finish_reason = c0.get("finish_reason")
|
||||
if finish_reason is not None:
|
||||
stop_reason = self._FINISH_REASON_TO_STOP.get(str(finish_reason), StopReason.UNKNOWN)
|
||||
# 先补齐 content_block_stop(thinking + text),再发送 MessageStop
|
||||
# 先补齐 content_block_stop(thinking + text + tool_calls),再发送 MessageStop
|
||||
if ss.get("thinking_block_started") and not ss.get("thinking_block_stopped"):
|
||||
ss["thinking_block_stopped"] = True
|
||||
events.append(
|
||||
@@ -660,6 +672,14 @@ class OpenAINormalizer(FormatNormalizer):
|
||||
events.append(
|
||||
ContentBlockStopEvent(block_index=_reserve_block_index("text_block_index"))
|
||||
)
|
||||
# 补齐所有已开始但未结束的 tool_call block
|
||||
tool_id_map = ss.get("tool_id_to_block_index")
|
||||
if isinstance(tool_id_map, dict):
|
||||
for _tk, bi in tool_id_map.items():
|
||||
stopped_key = f"tool_stopped:{bi}"
|
||||
if ss.get(f"tool_started:{bi}") and not ss.get(stopped_key):
|
||||
ss[stopped_key] = True
|
||||
events.append(ContentBlockStopEvent(block_index=int(bi)))
|
||||
# 解析 usage(需要请求时设置 stream_options.include_usage: true)
|
||||
usage_info = self._openai_usage_to_internal(chunk.get("usage"))
|
||||
events.append(MessageStopEvent(stop_reason=stop_reason, usage=usage_info))
|
||||
|
||||
@@ -55,6 +55,31 @@ from src.core.api_format.conversion.stream_events import (
|
||||
UnknownStreamEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
def _is_chat_completions_response(data: dict[str, Any]) -> bool:
|
||||
"""检测数据是否为 OpenAI Chat Completions 格式(而非 Responses API 格式)。
|
||||
|
||||
Chat Completions 的特征:
|
||||
- 非流式:有 choices 数组且 object == "chat.completion"
|
||||
- 流式:有 choices 数组且 object == "chat.completion.chunk"
|
||||
"""
|
||||
if not isinstance(data, dict):
|
||||
return False
|
||||
obj = data.get("object", "")
|
||||
if isinstance(obj, str) and obj.startswith("chat.completion"):
|
||||
return True
|
||||
if isinstance(data.get("choices"), list) and "type" not in data:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _get_openai_chat_normalizer() -> "FormatNormalizer | None":
|
||||
"""获取已注册的 openai:chat normalizer 实例(延迟获取避免循环导入)。"""
|
||||
from src.core.api_format.conversion.registry import format_conversion_registry
|
||||
|
||||
return format_conversion_registry.get_normalizer("openai:chat")
|
||||
|
||||
|
||||
class OpenAICliNormalizer(FormatNormalizer):
|
||||
@@ -126,7 +151,9 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
*,
|
||||
target_variant: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
is_codex = str(target_variant or "").lower() == "codex"
|
||||
openai_cli_extra = internal.extra.get("openai_cli", {})
|
||||
is_compact = bool(openai_cli_extra.get("_aether_compact"))
|
||||
is_codex = str(target_variant or "").lower() == "codex" and not is_compact
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"model": internal.model,
|
||||
@@ -178,7 +205,6 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
result["tool_choice"] = self._tool_choice_to_openai(internal.tool_choice)
|
||||
|
||||
# 还原 OpenAI Responses API 的其他字段(黑名单:已单独处理的字段不还原)
|
||||
openai_cli_extra = internal.extra.get("openai_cli", {})
|
||||
handled_keys = {
|
||||
"model",
|
||||
"input",
|
||||
@@ -228,16 +254,25 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
def response_to_internal(self, response: dict[str, Any]) -> InternalResponse:
|
||||
payload = self._unwrap_response_object(response)
|
||||
|
||||
# 检测 Chat Completions 格式回退
|
||||
if _is_chat_completions_response(payload):
|
||||
chat_norm = _get_openai_chat_normalizer()
|
||||
if chat_norm is not None:
|
||||
logger.debug(
|
||||
"[OpenAICliNormalizer] 检测到 Chat Completions 响应格式,委托给 openai:chat normalizer"
|
||||
)
|
||||
return chat_norm.response_to_internal(payload)
|
||||
|
||||
rid = str(payload.get("id") or "")
|
||||
model = str(payload.get("model") or "")
|
||||
|
||||
blocks, extra = self._extract_output_text_blocks(payload)
|
||||
blocks, extra, has_tool_use = self._extract_output_blocks(payload)
|
||||
usage = self._usage_to_internal(payload.get("usage"))
|
||||
|
||||
stop_reason = StopReason.UNKNOWN
|
||||
status = payload.get("status")
|
||||
if isinstance(status, str) and status == "completed":
|
||||
stop_reason = StopReason.END_TURN
|
||||
stop_reason = StopReason.TOOL_USE if has_tool_use else StopReason.END_TURN
|
||||
|
||||
return InternalResponse(
|
||||
id=rid,
|
||||
@@ -254,14 +289,49 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
*,
|
||||
requested_model: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
text = self._collapse_internal_text(internal.content)
|
||||
output_items: list[dict[str, Any]] = []
|
||||
|
||||
output_message = {
|
||||
"type": "message",
|
||||
"id": f"msg_{internal.id or 'stream'}",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": text}],
|
||||
}
|
||||
# 构建 output items:message(文本)和 function_call(工具调用)
|
||||
text = self._collapse_internal_text(internal.content)
|
||||
if text:
|
||||
output_items.append(
|
||||
{
|
||||
"type": "message",
|
||||
"id": f"msg_{internal.id or 'stream'}",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": text}],
|
||||
}
|
||||
)
|
||||
|
||||
for block in internal.content:
|
||||
if isinstance(block, ToolUseBlock):
|
||||
output_items.append(
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": block.tool_id,
|
||||
"id": block.tool_id,
|
||||
"name": block.tool_name,
|
||||
"arguments": (
|
||||
json.dumps(block.tool_input, ensure_ascii=False)
|
||||
if block.tool_input
|
||||
else "{}"
|
||||
),
|
||||
"status": "completed",
|
||||
}
|
||||
)
|
||||
|
||||
# 如果没有任何 output item,添加空 message(保持结构完整)
|
||||
if not output_items:
|
||||
output_items.append(
|
||||
{
|
||||
"type": "message",
|
||||
"id": f"msg_{internal.id or 'stream'}",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": ""}],
|
||||
}
|
||||
)
|
||||
|
||||
usage = internal.usage or UsageInfo()
|
||||
usage_obj: dict[str, Any] = {
|
||||
@@ -279,7 +349,7 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
"created": int(time.time()),
|
||||
"model": model_name,
|
||||
"status": "completed",
|
||||
"output": [output_message],
|
||||
"output": output_items,
|
||||
"usage": usage_obj,
|
||||
}
|
||||
|
||||
@@ -303,6 +373,18 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
pass
|
||||
return events
|
||||
|
||||
# 检测 Chat Completions 流式格式回退
|
||||
# 某些 Provider 即使配置为 openai:cli 也可能返回 Chat Completions 格式
|
||||
if _is_chat_completions_response(chunk):
|
||||
chat_norm = _get_openai_chat_normalizer()
|
||||
if chat_norm is not None:
|
||||
if not ss.get("_chat_fallback_logged"):
|
||||
ss["_chat_fallback_logged"] = True
|
||||
logger.debug(
|
||||
"[OpenAICliNormalizer] 检测到 Chat Completions 流式格式,委托给 openai:chat"
|
||||
)
|
||||
return chat_norm.stream_chunk_to_internal(chunk, state)
|
||||
|
||||
etype = str(chunk.get("type") or "")
|
||||
|
||||
# 尽量在首次事件补齐 message_start
|
||||
@@ -376,7 +458,17 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
ss["text_block_stopped"] = True
|
||||
events.append(ContentBlockStopEvent(block_index=0))
|
||||
|
||||
events.append(MessageStopEvent(stop_reason=StopReason.END_TURN, usage=usage))
|
||||
# 补齐所有已开始但未结束的 tool_call block
|
||||
active_tools = ss.get("active_tool_blocks")
|
||||
if isinstance(active_tools, dict):
|
||||
for tool_id, bi in list(active_tools.items()):
|
||||
events.append(ContentBlockStopEvent(block_index=bi))
|
||||
active_tools.clear()
|
||||
|
||||
# 根据流中是否出现过工具调用来判断 stop_reason
|
||||
has_tool_calls = bool(ss.get("tool_calls"))
|
||||
stop_reason = StopReason.TOOL_USE if has_tool_calls else StopReason.END_TURN
|
||||
events.append(MessageStopEvent(stop_reason=stop_reason, usage=usage))
|
||||
return events
|
||||
|
||||
def _handle_response_failed(
|
||||
@@ -411,21 +503,29 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type == "function_call":
|
||||
if not ss.get("tool_block_started"):
|
||||
ss["tool_block_started"] = True
|
||||
ss["current_tool_id"] = item.get("call_id") or item.get("id") or ""
|
||||
ss["current_tool_name"] = item.get("name") or ""
|
||||
events.append(
|
||||
ContentBlockStartEvent(
|
||||
block_index=ss.get("block_index", 0),
|
||||
block_type=ContentType.TOOL_USE,
|
||||
extra={
|
||||
"tool_id": ss["current_tool_id"],
|
||||
"tool_name": ss["current_tool_name"],
|
||||
},
|
||||
)
|
||||
tool_id = str(item.get("call_id") or item.get("id") or "")
|
||||
tool_name = str(item.get("name") or "")
|
||||
block_index = int(ss.get("block_index", 0))
|
||||
|
||||
# 记录当前活跃的工具调用(支持并行)
|
||||
active_tools = ss.setdefault("active_tool_blocks", {})
|
||||
active_tools[tool_id] = block_index
|
||||
ss["current_tool_id"] = tool_id
|
||||
ss["current_tool_name"] = tool_name
|
||||
|
||||
# 初始化工具调用收集
|
||||
tool_calls = ss.setdefault("tool_calls", {})
|
||||
tool_calls.setdefault(tool_id, {"name": tool_name, "args": ""})
|
||||
|
||||
events.append(
|
||||
ContentBlockStartEvent(
|
||||
block_index=block_index,
|
||||
block_type=ContentType.TOOL_USE,
|
||||
tool_id=tool_id,
|
||||
tool_name=tool_name,
|
||||
)
|
||||
ss["block_index"] = ss.get("block_index", 0) + 1
|
||||
)
|
||||
ss["block_index"] = block_index + 1
|
||||
return events
|
||||
|
||||
def _handle_output_item_done(
|
||||
@@ -435,9 +535,11 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
item = chunk.get("item")
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
if item_type == "function_call" and ss.get("tool_block_started"):
|
||||
ss["tool_block_started"] = False
|
||||
events.append(ContentBlockStopEvent(block_index=ss.get("block_index", 1) - 1))
|
||||
if item_type == "function_call":
|
||||
tool_id = str(item.get("call_id") or item.get("id") or "")
|
||||
active_tools = ss.get("active_tool_blocks", {})
|
||||
block_index = active_tools.pop(tool_id, ss.get("block_index", 1) - 1)
|
||||
events.append(ContentBlockStopEvent(block_index=block_index))
|
||||
return events
|
||||
|
||||
def _handle_function_call_delta(
|
||||
@@ -446,10 +548,20 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
events: list[InternalStreamEvent] = []
|
||||
delta = chunk.get("delta") or ""
|
||||
if delta:
|
||||
# 确定当前工具调用的 block_index 和 tool_id
|
||||
tool_id = str(chunk.get("item_id") or ss.get("current_tool_id", ""))
|
||||
active_tools = ss.get("active_tool_blocks", {})
|
||||
block_index = active_tools.get(tool_id, ss.get("block_index", 1) - 1)
|
||||
|
||||
# 累积参数
|
||||
tool_calls = ss.setdefault("tool_calls", {})
|
||||
entry = tool_calls.setdefault(tool_id, {"name": "", "args": ""})
|
||||
entry["args"] = str(entry.get("args") or "") + delta
|
||||
|
||||
events.append(
|
||||
ToolCallDeltaEvent(
|
||||
block_index=ss.get("block_index", 1) - 1,
|
||||
tool_id=ss.get("current_tool_id", ""),
|
||||
block_index=block_index,
|
||||
tool_id=tool_id,
|
||||
input_delta=delta,
|
||||
)
|
||||
)
|
||||
@@ -904,17 +1016,27 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
return resp_inner
|
||||
return response
|
||||
|
||||
def _extract_output_text_blocks(
|
||||
def _extract_output_blocks(
|
||||
self, payload: dict[str, Any]
|
||||
) -> tuple[list[ContentBlock], dict[str, Any]]:
|
||||
) -> tuple[list[ContentBlock], dict[str, Any], bool]:
|
||||
"""从 Responses API 的 output 提取所有内容块。
|
||||
|
||||
Returns:
|
||||
(blocks, extra, has_tool_use): 内容块列表、extra 信息、是否包含工具调用
|
||||
"""
|
||||
text_parts: list[str] = []
|
||||
blocks: list[ContentBlock] = []
|
||||
has_tool_use = False
|
||||
|
||||
output = payload.get("output")
|
||||
if isinstance(output, list):
|
||||
for item in output:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("type") == "message":
|
||||
|
||||
item_type = item.get("type")
|
||||
|
||||
if item_type == "message":
|
||||
content = item.get("content")
|
||||
if isinstance(content, list):
|
||||
for part in content:
|
||||
@@ -927,22 +1049,44 @@ class OpenAICliNormalizer(FormatNormalizer):
|
||||
text_parts.append(part.get("text") or "")
|
||||
continue
|
||||
|
||||
if item.get("type") in ("output_text", "text") and isinstance(
|
||||
item.get("text"), str
|
||||
):
|
||||
if item_type == "function_call":
|
||||
has_tool_use = True
|
||||
tool_id = str(item.get("call_id") or item.get("id") or "")
|
||||
tool_name = str(item.get("name") or "")
|
||||
args_raw = item.get("arguments") or "{}"
|
||||
try:
|
||||
tool_input = (
|
||||
json.loads(args_raw)
|
||||
if isinstance(args_raw, str)
|
||||
else (args_raw if isinstance(args_raw, dict) else {})
|
||||
)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
tool_input = {"_raw": args_raw}
|
||||
blocks.append(
|
||||
ToolUseBlock(
|
||||
tool_id=tool_id,
|
||||
tool_name=tool_name,
|
||||
tool_input=tool_input,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
if item_type in ("output_text", "text") and isinstance(item.get("text"), str):
|
||||
text_parts.append(item.get("text") or "")
|
||||
|
||||
# 兼容:部分实现可能直接给 output_text
|
||||
if not text_parts and isinstance(payload.get("output_text"), str):
|
||||
text_parts.append(payload.get("output_text") or "")
|
||||
|
||||
blocks: list[ContentBlock] = []
|
||||
# 文本块放在前面,工具调用块在后面(与 Claude 的 content 顺序一致)
|
||||
result_blocks: list[ContentBlock] = []
|
||||
text = "".join(text_parts)
|
||||
if text:
|
||||
blocks.append(TextBlock(text=text))
|
||||
result_blocks.append(TextBlock(text=text))
|
||||
result_blocks.extend(blocks)
|
||||
|
||||
extra: dict[str, Any] = {"raw": {"openai_cli_output": output}} if output is not None else {}
|
||||
return blocks, extra
|
||||
return result_blocks, extra, has_tool_use
|
||||
|
||||
def _usage_to_internal(self, usage: Any) -> UsageInfo:
|
||||
if not isinstance(usage, dict):
|
||||
|
||||
@@ -438,7 +438,7 @@ _REGISTRATION_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def register_default_normalizers() -> None:
|
||||
"""注册默认 Normalizers(OPENAI/CLAUDE/GEMINI + *_CLI)"""
|
||||
"""自动发现并注册 normalizers/ 目录下的所有 FormatNormalizer 实现"""
|
||||
global _DEFAULT_NORMALIZERS_REGISTERED # noqa: PLW0603 - module-level 缓存
|
||||
|
||||
# 快速路径:已注册则直接返回(无锁)
|
||||
@@ -450,23 +450,44 @@ def register_default_normalizers() -> None:
|
||||
if _DEFAULT_NORMALIZERS_REGISTERED:
|
||||
return
|
||||
|
||||
from src.core.api_format.conversion.normalizers.claude import ClaudeNormalizer
|
||||
from src.core.api_format.conversion.normalizers.claude_cli import ClaudeCliNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini import GeminiNormalizer
|
||||
from src.core.api_format.conversion.normalizers.gemini_cli import GeminiCliNormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai import OpenAINormalizer
|
||||
from src.core.api_format.conversion.normalizers.openai_cli import OpenAICliNormalizer
|
||||
import importlib
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
|
||||
format_conversion_registry.register(OpenAINormalizer())
|
||||
format_conversion_registry.register(OpenAICliNormalizer())
|
||||
format_conversion_registry.register(ClaudeNormalizer())
|
||||
format_conversion_registry.register(ClaudeCliNormalizer())
|
||||
format_conversion_registry.register(GeminiNormalizer())
|
||||
format_conversion_registry.register(GeminiCliNormalizer())
|
||||
normalizers_dir = Path(__file__).parent / "normalizers"
|
||||
for py_file in sorted(normalizers_dir.glob("*.py")):
|
||||
if py_file.name.startswith("_"):
|
||||
continue
|
||||
module_name = py_file.stem
|
||||
module_path = f"src.core.api_format.conversion.normalizers.{module_name}"
|
||||
try:
|
||||
mod = importlib.import_module(module_path)
|
||||
except Exception as e:
|
||||
logger.error("[FormatConversionRegistry] 导入 {} 失败: {}", module_path, e)
|
||||
continue
|
||||
for _attr_name, obj in inspect.getmembers(mod, inspect.isclass):
|
||||
if (
|
||||
issubclass(obj, FormatNormalizer)
|
||||
and obj is not FormatNormalizer
|
||||
and hasattr(obj, "FORMAT_ID")
|
||||
and obj.__module__ == mod.__name__
|
||||
):
|
||||
fmt_id = str(obj.FORMAT_ID).upper()
|
||||
if format_conversion_registry.get_normalizer(fmt_id) is not None:
|
||||
logger.warning(
|
||||
"[FormatConversionRegistry] FORMAT_ID '{}' 重复注册,{} 将覆盖已有实现",
|
||||
fmt_id,
|
||||
obj.__name__,
|
||||
)
|
||||
try:
|
||||
format_conversion_registry.register(obj())
|
||||
except Exception as e:
|
||||
logger.error("[FormatConversionRegistry] 注册 {} 失败: {}", obj.__name__, e)
|
||||
|
||||
_DEFAULT_NORMALIZERS_REGISTERED = True
|
||||
logger.info(
|
||||
f"[FormatConversionRegistry] 已注册 {len(format_conversion_registry.list_normalizers())} 个 normalizer"
|
||||
"[FormatConversionRegistry] 已注册 {} 个 normalizer",
|
||||
len(format_conversion_registry.list_normalizers()),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -334,19 +334,21 @@ class HeaderBuilder:
|
||||
def build(self) -> dict[str, str]:
|
||||
"""构建最终的头部字典
|
||||
|
||||
Safety net: 跳过值中包含非 ASCII 字符的头部并记录警告,
|
||||
防止 httpx 发送时抛出 ``UnicodeEncodeError``。
|
||||
httpx 要求 header 值可被 latin-1 编码。对于包含非 latin-1 字符
|
||||
(如中文)的值,先 UTF-8 编码再按 latin-1 解码,使 httpx 将原始
|
||||
UTF-8 字节逐字节发送到上游 —— 与 Go net/http 的行为一致。
|
||||
"""
|
||||
result: dict[str, str] = {}
|
||||
for original_key, value in self._headers.values():
|
||||
try:
|
||||
value.encode("ascii")
|
||||
value.encode("latin-1")
|
||||
except (UnicodeEncodeError, UnicodeDecodeError):
|
||||
logger.warning(
|
||||
"Dropping non-ASCII header before upstream request: {}",
|
||||
# 将 UTF-8 字节逐字节映射为 latin-1 字符串,httpx 会原样发送
|
||||
logger.debug(
|
||||
"Header '{}' contains non-latin-1 chars, encoding as raw UTF-8 bytes",
|
||||
original_key,
|
||||
)
|
||||
continue
|
||||
value = value.encode("utf-8").decode("latin-1")
|
||||
result[original_key] = value
|
||||
return result
|
||||
|
||||
|
||||
@@ -14,6 +14,17 @@ from src.core.modules.base import (
|
||||
ModuleMetadata,
|
||||
ModuleStatus,
|
||||
)
|
||||
from src.core.modules.hooks import (
|
||||
AUTH_AUTHENTICATE,
|
||||
AUTH_CHECK_EXCLUSIVE_MODE,
|
||||
AUTH_CHECK_REGISTRATION,
|
||||
AUTH_GET_METHODS,
|
||||
AUTH_TOKEN_PREFIX_AUTHENTICATORS,
|
||||
HookDispatcher,
|
||||
HookSpec,
|
||||
HookStrategy,
|
||||
get_hook_dispatcher,
|
||||
)
|
||||
from src.core.modules.registry import ModuleRegistry, get_module_registry
|
||||
|
||||
__all__ = [
|
||||
@@ -23,4 +34,14 @@ __all__ = [
|
||||
"ModuleStatus",
|
||||
"ModuleRegistry",
|
||||
"get_module_registry",
|
||||
# Hook system
|
||||
"HookDispatcher",
|
||||
"HookSpec",
|
||||
"HookStrategy",
|
||||
"get_hook_dispatcher",
|
||||
"AUTH_GET_METHODS",
|
||||
"AUTH_AUTHENTICATE",
|
||||
"AUTH_CHECK_REGISTRATION",
|
||||
"AUTH_CHECK_EXCLUSIVE_MODE",
|
||||
"AUTH_TOKEN_PREFIX_AUTHENTICATORS",
|
||||
]
|
||||
|
||||
@@ -91,6 +91,11 @@ class ModuleDefinition:
|
||||
# 配置验证(可选,启用模块时调用,返回 (success, error_message))
|
||||
validate_config: Callable[[Session], tuple[bool, str]] | None = None
|
||||
|
||||
# 钩子实现(可选)
|
||||
# {hook_name: handler_callable}
|
||||
# 模块通过此字段声明自己对核心扩展点的实现
|
||||
hooks: dict[str, Callable[..., Any]] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModuleStatus:
|
||||
|
||||
245
src/core/modules/hooks.py
Normal file
245
src/core/modules/hooks.py
Normal file
@@ -0,0 +1,245 @@
|
||||
"""
|
||||
模块钩子系统
|
||||
|
||||
提供模块与核心代码之间的动态扩展点。
|
||||
模块通过 ModuleDefinition.hooks 声明钩子实现,
|
||||
核心代码通过 HookDispatcher 调用所有活跃模块的钩子。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from inspect import isawaitable
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
# 钩子处理器类型: 可以是同步或异步函数
|
||||
HookHandler = Any # Callable[..., Any]
|
||||
|
||||
|
||||
class HookStrategy(str, Enum):
|
||||
"""钩子执行策略"""
|
||||
|
||||
FIRST_RESULT = "first_result" # 返回第一个非 None 结果
|
||||
COLLECT_ALL = "collect_all" # 收集所有结果到列表
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class HookSpec:
|
||||
"""钩子规格定义"""
|
||||
|
||||
name: str # 如 "auth.authenticate"
|
||||
strategy: HookStrategy = HookStrategy.FIRST_RESULT
|
||||
requires_active_check: bool = True # 是否过滤非活跃模块
|
||||
|
||||
|
||||
# ==================== 预定义钩子规格 ====================
|
||||
|
||||
AUTH_GET_METHODS = HookSpec(
|
||||
name="auth.get_methods",
|
||||
strategy=HookStrategy.COLLECT_ALL,
|
||||
)
|
||||
"""查询所有可用认证方法。返回 list[dict],每个 dict 包含认证方式信息。"""
|
||||
|
||||
AUTH_AUTHENTICATE = HookSpec(
|
||||
name="auth.authenticate",
|
||||
strategy=HookStrategy.FIRST_RESULT,
|
||||
)
|
||||
"""模块参与认证流程。kwargs: db, email, password, auth_type。返回 User 或 None。"""
|
||||
|
||||
AUTH_CHECK_REGISTRATION = HookSpec(
|
||||
name="auth.check_registration",
|
||||
strategy=HookStrategy.FIRST_RESULT,
|
||||
)
|
||||
"""模块检查是否允许本地注册。返回 {"blocked": True, "reason": "..."} 或 None。"""
|
||||
|
||||
AUTH_CHECK_EXCLUSIVE_MODE = HookSpec(
|
||||
name="auth.check_exclusive_mode",
|
||||
strategy=HookStrategy.FIRST_RESULT,
|
||||
)
|
||||
"""检查是否有模块开启了排他登录模式。返回 True 或 None。"""
|
||||
|
||||
AUTH_TOKEN_PREFIX_AUTHENTICATORS = HookSpec(
|
||||
name="auth.token_prefix_authenticators",
|
||||
strategy=HookStrategy.COLLECT_ALL,
|
||||
requires_active_check=False, # token 前缀认证是核心鉴权路径,只要模块已注册即可
|
||||
)
|
||||
"""获取 token 前缀认证器列表。返回 list[{"prefix": "ae_", "module": "..."}]。"""
|
||||
|
||||
|
||||
class HookDispatcher:
|
||||
"""
|
||||
钩子分发器 -- 单例
|
||||
|
||||
职责:
|
||||
- 注册模块的钩子实现
|
||||
- 在核心代码调用时,只执行活跃模块的钩子
|
||||
- 支持 FIRST_RESULT 和 COLLECT_ALL 两种执行策略
|
||||
"""
|
||||
|
||||
_instance: HookDispatcher | None = None
|
||||
|
||||
def __init__(self) -> None:
|
||||
# {hook_name: [(module_name, handler), ...]}
|
||||
self._handlers: defaultdict[str, list[tuple[str, HookHandler]]] = defaultdict(list)
|
||||
|
||||
@classmethod
|
||||
def get_instance(cls) -> HookDispatcher:
|
||||
if cls._instance is None:
|
||||
cls._instance = cls()
|
||||
return cls._instance
|
||||
|
||||
@classmethod
|
||||
def reset_instance(cls) -> None:
|
||||
"""重置单例(仅用于测试)"""
|
||||
cls._instance = None
|
||||
|
||||
def register(self, hook_name: str, module_name: str, handler: HookHandler) -> None:
|
||||
"""注册钩子处理器"""
|
||||
self._handlers[hook_name].append((module_name, handler))
|
||||
logger.debug("Hook [{}] registered handler from module [{}]", hook_name, module_name)
|
||||
|
||||
def has_handlers(self, hook_name: str) -> bool:
|
||||
"""检查是否有注册的处理器"""
|
||||
return bool(self._handlers.get(hook_name))
|
||||
|
||||
def _get_active_handlers(
|
||||
self, spec: HookSpec, db: Session | None
|
||||
) -> list[tuple[str, HookHandler]]:
|
||||
"""获取活跃模块的处理器列表"""
|
||||
handlers = self._handlers.get(spec.name, [])
|
||||
if not handlers:
|
||||
return []
|
||||
|
||||
if not spec.requires_active_check or db is None:
|
||||
return handlers
|
||||
|
||||
from src.core.modules.registry import get_module_registry
|
||||
|
||||
registry = get_module_registry()
|
||||
return [(name, handler) for name, handler in handlers if registry.is_active(name, db)]
|
||||
|
||||
# ==================== 异步分发 ====================
|
||||
|
||||
async def dispatch(
|
||||
self,
|
||||
spec: HookSpec,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
异步分发钩子调用
|
||||
|
||||
从 kwargs 中提取 db 参数用于活跃性检查,所有 kwargs 原样传递给处理器。
|
||||
|
||||
Args:
|
||||
spec: 钩子规格
|
||||
**kwargs: 传递给处理器的参数(其中 db 同时用于活跃性检查)
|
||||
|
||||
Returns:
|
||||
FIRST_RESULT: 第一个非 None 结果,或 None
|
||||
COLLECT_ALL: 结果列表
|
||||
"""
|
||||
db = kwargs.get("db")
|
||||
active_handlers = self._get_active_handlers(spec, db)
|
||||
if not active_handlers:
|
||||
return [] if spec.strategy == HookStrategy.COLLECT_ALL else None
|
||||
|
||||
if spec.strategy == HookStrategy.FIRST_RESULT:
|
||||
return await self._dispatch_first_result(spec.name, active_handlers, **kwargs)
|
||||
elif spec.strategy == HookStrategy.COLLECT_ALL:
|
||||
return await self._dispatch_collect_all(spec.name, active_handlers, **kwargs)
|
||||
return None
|
||||
|
||||
async def _call_handler(self, handler: HookHandler, **kwargs: Any) -> Any:
|
||||
"""调用处理器(支持同步和异步)"""
|
||||
result = handler(**kwargs)
|
||||
if isawaitable(result):
|
||||
return await result
|
||||
return result
|
||||
|
||||
async def _dispatch_first_result(
|
||||
self, hook_name: str, handlers: list[tuple[str, HookHandler]], **kwargs: Any
|
||||
) -> Any:
|
||||
for module_name, handler in handlers:
|
||||
try:
|
||||
result = await self._call_handler(handler, **kwargs)
|
||||
if result is not None:
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error("Hook [{}] handler from [{}] failed: {}", hook_name, module_name, e)
|
||||
return None
|
||||
|
||||
async def _dispatch_collect_all(
|
||||
self, hook_name: str, handlers: list[tuple[str, HookHandler]], **kwargs: Any
|
||||
) -> list[Any]:
|
||||
results: list[Any] = []
|
||||
for module_name, handler in handlers:
|
||||
try:
|
||||
result = await self._call_handler(handler, **kwargs)
|
||||
if result is not None:
|
||||
if isinstance(result, list):
|
||||
results.extend(result)
|
||||
else:
|
||||
results.append(result)
|
||||
except Exception as e:
|
||||
logger.error("Hook [{}] handler from [{}] failed: {}", hook_name, module_name, e)
|
||||
return results
|
||||
|
||||
# ==================== 同步分发 ====================
|
||||
|
||||
def dispatch_sync(
|
||||
self,
|
||||
spec: HookSpec,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
同步版本的 dispatch(仅适用于同步钩子处理器)
|
||||
|
||||
从 kwargs 中提取 db 参数用于活跃性检查,所有 kwargs 原样传递给处理器。
|
||||
用于无法使用 await 的同步上下文(如 OAuthService 的某些方法)。
|
||||
"""
|
||||
db = kwargs.get("db")
|
||||
active_handlers = self._get_active_handlers(spec, db)
|
||||
if not active_handlers:
|
||||
return [] if spec.strategy == HookStrategy.COLLECT_ALL else None
|
||||
|
||||
if spec.strategy == HookStrategy.FIRST_RESULT:
|
||||
for module_name, handler in active_handlers:
|
||||
try:
|
||||
result = handler(**kwargs)
|
||||
if result is not None:
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Hook [{}] sync handler from [{}] failed: {}", spec.name, module_name, e
|
||||
)
|
||||
return None
|
||||
|
||||
elif spec.strategy == HookStrategy.COLLECT_ALL:
|
||||
results: list[Any] = []
|
||||
for module_name, handler in active_handlers:
|
||||
try:
|
||||
result = handler(**kwargs)
|
||||
if result is not None:
|
||||
if isinstance(result, list):
|
||||
results.extend(result)
|
||||
else:
|
||||
results.append(result)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Hook [{}] sync handler from [{}] failed: {}", spec.name, module_name, e
|
||||
)
|
||||
return results
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_hook_dispatcher() -> HookDispatcher:
|
||||
"""获取钩子分发器实例"""
|
||||
return HookDispatcher.get_instance()
|
||||
Reference in New Issue
Block a user