|
22 | 22 |
|
23 | 23 | from __future__ import annotations |
24 | 24 |
|
| 25 | +import contextlib |
25 | 26 | import contextvars |
| 27 | +import os |
26 | 28 | from dataclasses import dataclass |
27 | | -from typing import Optional |
| 29 | +from typing import Iterator, Mapping, Optional |
| 30 | + |
| 31 | +# FC 注入 STS 的默认头名(可经构造参数或环境变量覆盖)。 |
| 32 | +DEFAULT_ACCESS_KEY_ID_HEADER = "x-fc-access-key-id" |
| 33 | +DEFAULT_ACCESS_KEY_SECRET_HEADER = "x-fc-access-key-secret" |
| 34 | +DEFAULT_SECURITY_TOKEN_HEADER = "x-fc-security-token" |
28 | 35 |
|
29 | 36 |
|
30 | 37 | @dataclass(frozen=True) |
@@ -78,3 +85,109 @@ def reset_request_sts(token: contextvars.Token) -> None: |
78 | 85 | def get_request_sts() -> Optional[StsCredential]: |
79 | 86 | """获取当前请求的 STS 覆盖层;未设置时返回 ``None``。""" |
80 | 87 | return _current_sts.get() |
| 88 | + |
| 89 | + |
| 90 | +def _resolve_header_name( |
| 91 | + explicit: Optional[str], env_key: str, default: str |
| 92 | +) -> str: |
| 93 | + """解析头名:构造参数 > 环境变量 > 默认值;统一转小写。""" |
| 94 | + return (explicit or os.getenv(env_key) or default).lower() |
| 95 | + |
| 96 | + |
| 97 | +def sts_from_headers( |
| 98 | + headers: Mapping[str, str], |
| 99 | + *, |
| 100 | + access_key_id_header: Optional[str] = None, |
| 101 | + access_key_secret_header: Optional[str] = None, |
| 102 | + security_token_header: Optional[str] = None, |
| 103 | +) -> Optional[StsCredential]: |
| 104 | + """从请求头映射解析 STS 三元组;不齐全则返回 ``None``。 |
| 105 | +
|
| 106 | + 仅当 ak/sk/sts 三者齐全才视为有效刷新(避免把新 sts 与陈旧/环境变量里的 |
| 107 | + ak/sk 混用)。``headers`` 可为任意 Mapping(如 ``dict`` 或 Starlette |
| 108 | + ``Headers``),按头名**大小写不敏感**查找。头名优先级:参数 > 环境变量 |
| 109 | + (``AGENTRUN_STS_HEADER_*``)> 默认(``x-fc-*``)。 |
| 110 | + """ |
| 111 | + ak_name = _resolve_header_name( |
| 112 | + access_key_id_header, |
| 113 | + "AGENTRUN_STS_HEADER_ACCESS_KEY_ID", |
| 114 | + DEFAULT_ACCESS_KEY_ID_HEADER, |
| 115 | + ) |
| 116 | + sk_name = _resolve_header_name( |
| 117 | + access_key_secret_header, |
| 118 | + "AGENTRUN_STS_HEADER_ACCESS_KEY_SECRET", |
| 119 | + DEFAULT_ACCESS_KEY_SECRET_HEADER, |
| 120 | + ) |
| 121 | + sts_name = _resolve_header_name( |
| 122 | + security_token_header, |
| 123 | + "AGENTRUN_STS_HEADER_SECURITY_TOKEN", |
| 124 | + DEFAULT_SECURITY_TOKEN_HEADER, |
| 125 | + ) |
| 126 | + |
| 127 | + lower = {str(k).lower(): v for k, v in headers.items()} |
| 128 | + cred = StsCredential( |
| 129 | + access_key_id=lower.get(ak_name), |
| 130 | + access_key_secret=lower.get(sk_name), |
| 131 | + security_token=lower.get(sts_name), |
| 132 | + ) |
| 133 | + return cred if cred.is_complete() else None |
| 134 | + |
| 135 | + |
| 136 | +@contextlib.contextmanager |
| 137 | +def use_sts_credentials( |
| 138 | + access_key_id: Optional[str] = None, |
| 139 | + access_key_secret: Optional[str] = None, |
| 140 | + security_token: Optional[str] = None, |
| 141 | +) -> Iterator[StsCredential]: |
| 142 | + """在 ``with`` 块内临时使用给定 STS 临时凭证(请求级 overlay),退出自动复位。 |
| 143 | +
|
| 144 | + 适用于**不经过 agentrun server** 的场景:自有 FastAPI / Flask / Django,或 |
| 145 | + 非 HTTP 的任务里,从上游 / 请求头拿到最新 STS 后注入——块内所有 SDK 调用 |
| 146 | + (以及其内创建的 asyncio 任务)即使用这组凭证。 |
| 147 | +
|
| 148 | + Examples: |
| 149 | + >>> with use_sts_credentials(ak, sk, sts): |
| 150 | + ... knowledgebase.retrieve(...) # 使用最新 STS |
| 151 | +
|
| 152 | + Note: |
| 153 | + 基于 ``contextvars``,按当前任务/线程隔离;用户自行 ``threading.Thread`` |
| 154 | + 起的裸线程不会继承(``asyncio.create_task`` 会)。 |
| 155 | + """ |
| 156 | + cred = StsCredential(access_key_id, access_key_secret, security_token) |
| 157 | + token = set_request_sts(cred) |
| 158 | + try: |
| 159 | + yield cred |
| 160 | + finally: |
| 161 | + reset_request_sts(token) |
| 162 | + |
| 163 | + |
| 164 | +@contextlib.contextmanager |
| 165 | +def use_sts_from_headers( |
| 166 | + headers: Mapping[str, str], |
| 167 | + *, |
| 168 | + access_key_id_header: Optional[str] = None, |
| 169 | + access_key_secret_header: Optional[str] = None, |
| 170 | + security_token_header: Optional[str] = None, |
| 171 | +) -> Iterator[Optional[StsCredential]]: |
| 172 | + """从请求头映射解析 STS 并在 ``with`` 块内生效;三元组不齐全则不覆盖(透传)。 |
| 173 | +
|
| 174 | + 与 :class:`agentrun.server.sts_middleware.StsRefreshMiddleware` 共用同一套 |
| 175 | + 解析逻辑。适用于在自有 Web 框架里手动接入: |
| 176 | +
|
| 177 | + >>> with use_sts_from_headers(request.headers): |
| 178 | + ... await invoke_agent(request) |
| 179 | + """ |
| 180 | + cred = sts_from_headers( |
| 181 | + headers, |
| 182 | + access_key_id_header=access_key_id_header, |
| 183 | + access_key_secret_header=access_key_secret_header, |
| 184 | + security_token_header=security_token_header, |
| 185 | + ) |
| 186 | + if cred is None: |
| 187 | + yield None |
| 188 | + return |
| 189 | + token = set_request_sts(cred) |
| 190 | + try: |
| 191 | + yield cred |
| 192 | + finally: |
| 193 | + reset_request_sts(token) |
0 commit comments