mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
feat(codex): 引入 upstream_headers hook 机制,为 Codex 注入 session/conversation/account headers
- 新增 upstream_headers.py:可注册 provider+endpoint 维度的 extra headers 构建 hook - Codex openai:cli 注入 session_id + conversation_id(由 prompt_cache_key sha256 派生) - Codex openai:compact 注入 chatgpt-account-id(来自 auth_config)+ session_id,不注入 conversation_id - 修复 prompt_cache:compact 格式现统一为 codex 策略,不再跳过注入 - chat_handler_base / cli_request_mixin 均在 extra_headers 阶段调用 build_upstream_extra_headers
This commit is contained in:
@@ -91,6 +91,7 @@ from src.services.provider.stream_policy import (
|
||||
from src.services.provider.transport import (
|
||||
build_provider_url,
|
||||
)
|
||||
from src.services.provider.upstream_headers import build_upstream_extra_headers
|
||||
from src.services.scheduling.aware_scheduler import ProviderCandidate
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.task.request_state import MutableRequestBodyState
|
||||
@@ -836,6 +837,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
hook_headers = build_upstream_extra_headers(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=str(provider_api_format) if provider_api_format else None,
|
||||
request_body=request_body,
|
||||
original_headers=original_headers,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
if hook_headers:
|
||||
extra_headers.update(hook_headers)
|
||||
|
||||
return ProviderRequestResult(
|
||||
request_body=request_body,
|
||||
url_model=url_model,
|
||||
|
||||
@@ -20,6 +20,7 @@ from src.services.provider.stream_policy import (
|
||||
resolve_upstream_is_stream,
|
||||
)
|
||||
from src.services.provider.transport import build_provider_url
|
||||
from src.services.provider.upstream_headers import build_upstream_extra_headers
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from src.api.handlers.base.cli_protocol import CliHandlerProtocol
|
||||
@@ -302,6 +303,16 @@ class CliRequestMixin:
|
||||
if envelope:
|
||||
extra_headers.update(envelope.extra_headers() or {})
|
||||
|
||||
hook_headers = build_upstream_extra_headers(
|
||||
provider_type=provider_type,
|
||||
endpoint_sig=provider_api_format,
|
||||
request_body=request_body,
|
||||
original_headers=original_headers,
|
||||
decrypted_auth_config=auth_info.decrypted_auth_config if auth_info else None,
|
||||
)
|
||||
if hook_headers:
|
||||
extra_headers.update(hook_headers)
|
||||
|
||||
provider_payload, provider_headers = self._request_builder.build(
|
||||
request_body,
|
||||
original_headers,
|
||||
|
||||
@@ -11,6 +11,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
@@ -31,6 +33,85 @@ fetch_models_codex = create_preset_models_fetcher("codex")
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_header_value(headers: Mapping[str, Any] | None, header_name: str) -> str | None:
|
||||
if not isinstance(headers, Mapping):
|
||||
return None
|
||||
|
||||
target = str(header_name or "").strip().lower()
|
||||
if not target:
|
||||
return None
|
||||
|
||||
for name, value in headers.items():
|
||||
if str(name or "").strip().lower() != target:
|
||||
continue
|
||||
normalized = str(value or "").strip()
|
||||
return normalized or None
|
||||
return None
|
||||
|
||||
|
||||
def _build_short_header_id(seed: str) -> str:
|
||||
return hashlib.sha256(seed.encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
def _build_codex_headers(
|
||||
request_body: Any,
|
||||
original_headers: Mapping[str, Any] | None,
|
||||
*,
|
||||
include_conversation_id: bool,
|
||||
decrypted_auth_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""Codex upstream: chatgpt-account-id + session_id + conversation_id."""
|
||||
headers: dict[str, str] = {}
|
||||
|
||||
if decrypted_auth_config:
|
||||
account_id = str(decrypted_auth_config.get("account_id") or "").strip()
|
||||
if account_id:
|
||||
headers["chatgpt-account-id"] = account_id
|
||||
|
||||
if isinstance(request_body, dict):
|
||||
cache_key = str(request_body.get("prompt_cache_key") or "").strip()
|
||||
if cache_key:
|
||||
short_id = _build_short_header_id(cache_key)
|
||||
if not _get_header_value(original_headers, "session_id"):
|
||||
headers["session_id"] = short_id
|
||||
if include_conversation_id and not _get_header_value(
|
||||
original_headers, "conversation_id"
|
||||
):
|
||||
headers["conversation_id"] = short_id
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
def build_codex_cli_headers(
|
||||
request_body: Any,
|
||||
original_headers: Mapping[str, Any] | None,
|
||||
*,
|
||||
decrypted_auth_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
from src.services.provider.adapters.codex.context import is_codex_compact_request
|
||||
|
||||
return _build_codex_headers(
|
||||
request_body,
|
||||
original_headers,
|
||||
include_conversation_id=not is_codex_compact_request(endpoint_sig="openai:cli"),
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
|
||||
|
||||
def build_codex_compact_headers(
|
||||
request_body: Any,
|
||||
original_headers: Mapping[str, Any] | None,
|
||||
*,
|
||||
decrypted_auth_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, str]:
|
||||
return _build_codex_headers(
|
||||
request_body,
|
||||
original_headers,
|
||||
include_conversation_id=False,
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
|
||||
|
||||
def build_codex_url(
|
||||
endpoint: Any,
|
||||
*,
|
||||
@@ -165,10 +246,13 @@ def register_all() -> None:
|
||||
from src.core.provider_oauth_utils import register_auth_enricher
|
||||
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
|
||||
from src.services.provider.transport import register_transport_hook
|
||||
from src.services.provider.upstream_headers import register_upstream_headers_hook
|
||||
|
||||
# Transport
|
||||
register_transport_hook("codex", "openai:cli", build_codex_url)
|
||||
register_transport_hook("codex", "openai:compact", build_codex_url)
|
||||
register_upstream_headers_hook("codex", "openai:cli", build_codex_cli_headers)
|
||||
register_upstream_headers_hook("codex", "openai:compact", build_codex_compact_headers)
|
||||
|
||||
# Auth
|
||||
register_auth_enricher("codex", enrich_codex)
|
||||
|
||||
@@ -8,7 +8,6 @@ from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from src.core.provider_types import ProviderType, normalize_provider_type
|
||||
from src.services.provider.adapters.codex.context import is_codex_compact_request
|
||||
from src.utils.url_utils import is_official_openai_api_url
|
||||
|
||||
_OFFICIAL_OPENAI_PROMPT_CACHE_FORMATS: frozenset[str] = frozenset({"openai:chat", "openai:cli"})
|
||||
@@ -112,15 +111,13 @@ def resolve_prompt_cache_key_scope(
|
||||
) -> str | None:
|
||||
"""Resolve which prompt cache strategy should be used for this request."""
|
||||
fmt = str(provider_api_format or "").strip().lower()
|
||||
pt = normalize_provider_type(provider_type)
|
||||
if pt == ProviderType.CODEX.value and fmt in {"openai:cli", "openai:compact"}:
|
||||
return "codex"
|
||||
|
||||
if fmt == "openai:compact":
|
||||
return None
|
||||
|
||||
pt = normalize_provider_type(provider_type)
|
||||
if pt == ProviderType.CODEX.value and fmt == "openai:cli":
|
||||
if is_codex_compact_request(endpoint_sig=fmt):
|
||||
return None
|
||||
return "codex"
|
||||
|
||||
if fmt in _OFFICIAL_OPENAI_PROMPT_CACHE_FORMATS and is_official_openai_api_url(base_url):
|
||||
return "openai"
|
||||
|
||||
|
||||
59
src/services/provider/upstream_headers.py
Normal file
59
src/services/provider/upstream_headers.py
Normal file
@@ -0,0 +1,59 @@
|
||||
"""Provider-specific upstream request header hooks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Callable
|
||||
|
||||
from src.core.provider_types import normalize_provider_type
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
UpstreamHeadersHookFn = Callable[..., dict[str, str]]
|
||||
|
||||
_hooks: dict[tuple[str, str], UpstreamHeadersHookFn] = {}
|
||||
|
||||
|
||||
def register_upstream_headers_hook(
|
||||
provider_type: str,
|
||||
endpoint_sig: str,
|
||||
hook: UpstreamHeadersHookFn,
|
||||
) -> None:
|
||||
"""Register a provider-specific extra upstream headers builder."""
|
||||
pt = normalize_provider_type(provider_type)
|
||||
sig = str(endpoint_sig or "").strip().lower()
|
||||
if not pt or not sig:
|
||||
return
|
||||
_hooks[(pt, sig)] = hook
|
||||
|
||||
|
||||
def build_upstream_extra_headers(
|
||||
*,
|
||||
provider_type: str | None,
|
||||
endpoint_sig: str | None,
|
||||
request_body: Any,
|
||||
original_headers: Mapping[str, Any] | None,
|
||||
decrypted_auth_config: dict[str, Any] | None,
|
||||
) -> dict[str, str]:
|
||||
"""Build provider-specific extra upstream headers for the current request."""
|
||||
pt = normalize_provider_type(provider_type)
|
||||
sig = str(endpoint_sig or "").strip().lower()
|
||||
if not pt or not sig:
|
||||
return {}
|
||||
|
||||
ensure_providers_bootstrapped(provider_types=[pt])
|
||||
hook = _hooks.get((pt, sig))
|
||||
if hook is None:
|
||||
return {}
|
||||
|
||||
return hook(
|
||||
request_body,
|
||||
original_headers,
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"UpstreamHeadersHookFn",
|
||||
"build_upstream_extra_headers",
|
||||
"register_upstream_headers_hook",
|
||||
]
|
||||
Reference in New Issue
Block a user