Skip to content

Commit 5e11488

Browse files
OhYeeclaude
andcommitted
feat(credential): 暴露 use_sts_credentials / use_sts_from_headers 供非 server 场景手动注入 STS
StsRefreshMiddleware 只在 agentrun server 内生效;用户用自有 Web 框架(FastAPI/ Flask/Django)或非 HTTP 任务时无法被拦截。新增两个公开上下文管理器(从 agentrun 顶层导出),让用户手动注入最新 STS: - use_sts_credentials(ak, sk, sts):显式传值 - use_sts_from_headers(headers):从任意请求头映射解析(同 x-fc-* 约定,大小写不敏感) 二者写入同一请求级 overlay,刷新逻辑与中间件完全一致(含三元组齐全校验、自动复位)。 中间件本身重构为直接复用 use_sts_from_headers —— 单一真相,server 内外同一套逻辑。 Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> Signed-off-by: OhYee <oyohyee@oyohyee.com>
1 parent d4073c6 commit 5e11488

5 files changed

Lines changed: 221 additions & 64 deletions

File tree

AGENTS.md

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -481,7 +481,19 @@ STS 临时凭证(ak/sk/security_token)会过期。部署在函数计算(FC
481481
1. **请求级 overlay**`agentrun/server/sts_middleware.py` 解析 FC 头
482482
(默认 `x-fc-access-key-id` / `x-fc-access-key-secret` / `x-fc-security-token`
483483
可经构造参数或 `AGENTRUN_STS_HEADER_*` 环境变量覆盖),写入
484-
`agentrun/utils/credential_context.py``contextvars` overlay。
484+
`agentrun/utils/credential_context.py``contextvars` overlay。中间件本身
485+
只是 `use_sts_from_headers` 的薄封装(加 FC 门控),二者共用同一套解析逻辑。
486+
487+
**非 agentrun server 场景**(自有 FastAPI / Flask / Django、或非 HTTP 任务):
488+
中间件不会运行,需用户手动注入。SDK 顶层导出两个上下文管理器:
489+
- `agentrun.use_sts_credentials(ak, sk, sts)` —— 显式传值;
490+
- `agentrun.use_sts_from_headers(headers)` —— 从任意请求头映射解析(同 `x-fc-*`)。
491+
492+
```python
493+
from agentrun import use_sts_from_headers
494+
with use_sts_from_headers(request.headers):
495+
... # 块内所有 SDK 调用使用最新 STS,退出自动复位
496+
```
485497

486498
2. **Config 懒解析**`Config` 的三个凭证 getter 按
487499
**显式传入 > 请求级 overlay(仅当三者均未显式传入)> 环境变量** 解析。

agentrun/__init__.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -135,6 +135,11 @@
135135
# ToolSet
136136
from agentrun.toolset import ToolSet, ToolSetClient
137137
from agentrun.utils.config import Config
138+
from agentrun.utils.credential_context import (
139+
StsCredential,
140+
use_sts_credentials,
141+
use_sts_from_headers,
142+
)
138143
from agentrun.utils.exception import (
139144
ResourceAlreadyExistError,
140145
ResourceNotExistError,
@@ -335,6 +340,10 @@
335340
"ResourceNotExistError",
336341
"ResourceAlreadyExistError",
337342
"Config",
343+
######## STS 凭证刷新(非 server 场景手动注入) ########
344+
"StsCredential",
345+
"use_sts_credentials",
346+
"use_sts_from_headers",
338347
]
339348

340349
# Memory Collection 模块的所有导出(延迟加载)

agentrun/server/sts_middleware.py

Lines changed: 13 additions & 62 deletions
Original file line numberDiff line numberDiff line change
@@ -33,16 +33,7 @@
3333
from starlette.responses import Response
3434
from starlette.types import ASGIApp
3535

36-
from agentrun.utils.credential_context import (
37-
StsCredential,
38-
reset_request_sts,
39-
set_request_sts,
40-
)
41-
from agentrun.utils.log import logger
42-
43-
DEFAULT_ACCESS_KEY_ID_HEADER = "x-fc-access-key-id"
44-
DEFAULT_ACCESS_KEY_SECRET_HEADER = "x-fc-access-key-secret"
45-
DEFAULT_SECURITY_TOKEN_HEADER = "x-fc-security-token"
36+
from agentrun.utils.credential_context import use_sts_from_headers
4637

4738

4839
def _detect_enabled() -> bool:
@@ -69,61 +60,21 @@ def __init__(
6960
super().__init__(app)
7061
# enabled=None 时自动探测(FC 环境或显式环境变量开关)。
7162
self._enabled = _detect_enabled() if enabled is None else enabled
72-
self._ak_header = (
73-
access_key_id_header
74-
or os.getenv("AGENTRUN_STS_HEADER_ACCESS_KEY_ID")
75-
or DEFAULT_ACCESS_KEY_ID_HEADER
76-
)
77-
self._sk_header = (
78-
access_key_secret_header
79-
or os.getenv("AGENTRUN_STS_HEADER_ACCESS_KEY_SECRET")
80-
or DEFAULT_ACCESS_KEY_SECRET_HEADER
81-
)
82-
self._sts_header = (
83-
security_token_header
84-
or os.getenv("AGENTRUN_STS_HEADER_SECURITY_TOKEN")
85-
or DEFAULT_SECURITY_TOKEN_HEADER
86-
)
63+
# 头名解析(参数 > 环境变量 > 默认)交由 sts_from_headers 处理,这里只存原值。
64+
self._ak_header = access_key_id_header
65+
self._sk_header = access_key_secret_header
66+
self._sts_header = security_token_header
8767

8868
async def dispatch(self, request: Request, call_next) -> Response:
8969
if not self._enabled:
9070
return await call_next(request)
9171

92-
cred = self._extract(request)
93-
if cred is None:
94-
return await call_next(request)
95-
96-
token = set_request_sts(cred)
97-
try:
72+
# 直接复用公开上下文管理器:解析请求头 -> 注入 overlay -> 退出复位。
73+
# 三元组不齐全时 use_sts_from_headers 不覆盖(透传),与手动注入完全一致。
74+
with use_sts_from_headers(
75+
request.headers,
76+
access_key_id_header=self._ak_header,
77+
access_key_secret_header=self._sk_header,
78+
security_token_header=self._sts_header,
79+
):
9880
return await call_next(request)
99-
finally:
100-
reset_request_sts(token)
101-
102-
def _extract(self, request: Request) -> Optional[StsCredential]:
103-
# starlette Headers 大小写不敏感。
104-
headers = request.headers
105-
cred = StsCredential(
106-
access_key_id=headers.get(self._ak_header),
107-
access_key_secret=headers.get(self._sk_header),
108-
security_token=headers.get(self._sts_header),
109-
)
110-
111-
# 只有三元组齐全才视为有效 STS 刷新,避免把新 sts 与陈旧 ak/sk 混用。
112-
if not cred.is_complete():
113-
if not (
114-
cred.access_key_id
115-
or cred.access_key_secret
116-
or cred.security_token
117-
):
118-
return None # 无任何凭证头:常规非 FC 请求,静默跳过。
119-
logger.warning(
120-
"STS headers incomplete (ak=%s, sk=%s, sts=%s); ignoring"
121-
" partial credential set and falling back to env.",
122-
"set" if cred.access_key_id else "unset",
123-
"set" if cred.access_key_secret else "unset",
124-
"set" if cred.security_token else "unset",
125-
)
126-
return None
127-
128-
logger.debug("STS overlay applied from request headers")
129-
return cred

agentrun/utils/credential_context.py

Lines changed: 114 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,9 +22,16 @@
2222

2323
from __future__ import annotations
2424

25+
import contextlib
2526
import contextvars
27+
import os
2628
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"
2835

2936

3037
@dataclass(frozen=True)
@@ -78,3 +85,109 @@ def reset_request_sts(token: contextvars.Token) -> None:
7885
def get_request_sts() -> Optional[StsCredential]:
7986
"""获取当前请求的 STS 覆盖层;未设置时返回 ``None``。"""
8087
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)

tests/unittests/test_sts_refresh.py

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,8 @@
1818
get_request_sts,
1919
reset_request_sts,
2020
set_request_sts,
21+
use_sts_credentials,
22+
use_sts_from_headers,
2123
)
2224

2325

@@ -418,3 +420,73 @@ async def run():
418420
if e.event == EventType.TEXT
419421
)
420422
assert text == "ak=OV_AK;sts=OV_STS", text
423+
424+
425+
# --------------------------------------------------------------------------- #
426+
# 公开 API:非 server 场景手动注入
427+
# --------------------------------------------------------------------------- #
428+
def test_use_sts_credentials_context_manager(monkeypatch):
429+
monkeypatch.setenv("AGENTRUN_ACCESS_KEY_ID", "ENV_AK")
430+
cfg = Config()
431+
assert cfg.get_access_key_id() == "ENV_AK"
432+
with use_sts_credentials("OV_AK", "OV_SK", "OV_STS"):
433+
assert cfg.get_access_key_id() == "OV_AK"
434+
assert cfg.get_access_key_secret() == "OV_SK"
435+
assert cfg.get_security_token() == "OV_STS"
436+
# 退出自动复位
437+
assert cfg.get_access_key_id() == "ENV_AK"
438+
assert get_request_sts() is None
439+
440+
441+
def test_use_sts_from_headers_complete_case_insensitive():
442+
cfg = Config()
443+
headers = {
444+
"X-Fc-Access-Key-Id": "H_AK", # 大小写不敏感
445+
"x-fc-access-key-secret": "H_SK",
446+
"X-FC-SECURITY-TOKEN": "H_STS",
447+
}
448+
with use_sts_from_headers(headers) as cred:
449+
assert cred is not None
450+
assert cfg.get_access_key_id() == "H_AK"
451+
assert cfg.get_security_token() == "H_STS"
452+
assert get_request_sts() is None
453+
454+
455+
def test_use_sts_from_headers_partial_no_override(monkeypatch):
456+
monkeypatch.setenv("AGENTRUN_ACCESS_KEY_ID", "ENV_AK")
457+
cfg = Config()
458+
# 只有 sts、缺 ak/sk -> 不构成完整三元组 -> 不覆盖
459+
with use_sts_from_headers({"x-fc-security-token": "H_STS"}) as cred:
460+
assert cred is None
461+
assert cfg.get_access_key_id() == "ENV_AK"
462+
assert cfg.get_security_token() == ""
463+
assert get_request_sts() is None
464+
465+
466+
def test_sts_from_headers_helper():
467+
from agentrun.utils.credential_context import sts_from_headers
468+
469+
assert sts_from_headers({"x-fc-access-key-id": "a"}) is None # 不齐全
470+
cred = sts_from_headers({
471+
"x-fc-access-key-id": "a",
472+
"x-fc-access-key-secret": "b",
473+
"x-fc-security-token": "c",
474+
})
475+
assert cred is not None
476+
assert (
477+
cred.access_key_id,
478+
cred.access_key_secret,
479+
cred.security_token,
480+
) == ("a", "b", "c")
481+
482+
483+
def test_public_exports_available():
484+
import agentrun
485+
486+
for name in (
487+
"StsCredential",
488+
"use_sts_credentials",
489+
"use_sts_from_headers",
490+
):
491+
assert hasattr(agentrun, name), f"{name} not exported"
492+
assert name in agentrun.__all__, f"{name} missing from __all__"

0 commit comments

Comments
 (0)