Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -38,3 +38,7 @@ deploy/docker/.venv-modelscope/
.idea/
.vscode/
CLAUDE.md

# 安全模块
# 开发子plan
/security-plans
21 changes: 17 additions & 4 deletions bootstrap/cli/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,14 +41,27 @@ def build_parser() -> argparse.ArgumentParser:
description="agent-memory memory engine CLI",
)
parser.add_argument(
"--server", "--base-url", dest="server",
metavar="URL", default=os.environ.get("AGENT_MEMORY_SERVER"),
"--server",
"--base-url",
dest="server",
metavar="URL",
default=os.environ.get("AGENT_MEMORY_SERVER"),
help="drive a running server over HTTP (Mem0 --base-url; default: in-process)",
)
parser.add_argument(
"--config", action="append", default=[], metavar="PATH",
"--config",
action="append",
default=[],
metavar="PATH",
help="JSON config layer stacked on OFFLINE (in-process only; repeatable)",
)
parser.add_argument(
"--api-key",
dest="api_key",
metavar="KEY",
default=None,
help="API key for --server mode (default: $AGENT_MEMORY_API_KEY)",
)

sub = parser.add_subparsers(dest="command", required=True)

Expand Down Expand Up @@ -78,7 +91,7 @@ def main(argv: list[str] | None = None) -> int:
sys.stderr.write("note: --config is ignored in --server (HTTP) mode\n")

try:
client = make_client(args.server, args.config)
client = make_client(args.server, args.config, args.api_key)
if args.command in ("health", "status"):
return commands.run_health(client, args)
if args.command == "batch":
Expand Down
41 changes: 35 additions & 6 deletions bootstrap/cli/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,9 +41,11 @@ class EngineClient(Protocol):
"""A backend the CLI can drive: turn a (verb, payload) into (status, body)."""

def call(self, verb: str, payload: dict[str, Any]) -> tuple[int, dict[str, Any]]:
"""Dispatch one memory-engine verb."""
...

def healthz(self) -> tuple[int, dict[str, Any]]:
"""Return the backend health response."""
...


Expand Down Expand Up @@ -74,9 +76,24 @@ def server(self):
return self._srv

def call(self, verb: str, payload: dict[str, Any]) -> tuple[int, dict[str, Any]]:
from auth_middleware import authenticated
from handler import dispatch

return dispatch(self._srv, verb, payload)
from common.errors import AuthenticationError
from security.types import Credentials

# 进程内直连没有 HTTP header,故过一个空 Credentials。DEV 模式下得到
# ROOT,与现状一致(CLI 一直是全权限的);API_KEY 模式下会认证失败——
# 这是**正确的**:没有凭据就不该有权限。要在 API_KEY 模式下用 CLI,
# 走 HttpClient 带 --api-key。
#
# 认证失败转成 (401, body) 而非抛出:本方法的契约是返回状态码,
# 与 HttpClient.call 一致。
try:
with authenticated(self._srv.authenticator, Credentials(), self._srv.audit):
return dispatch(self._srv, verb, payload)
except AuthenticationError as exc:
return 401, {"error": type(exc).__name__, "message": str(exc)}

def healthz(self) -> tuple[int, dict[str, Any]]:
return 200, {"status": "ok", "profile": self._srv.config.profile}
Expand All @@ -85,17 +102,21 @@ def healthz(self) -> tuple[int, dict[str, Any]]:
class HttpClient:
"""Drive a running ``bootstrap`` server over HTTP (``POST /v1/<verb>``)."""

def __init__(self, base_url: str, timeout: float = 30.0) -> None:
def __init__(self, base_url: str, timeout: float = 30.0, api_key: str = "") -> None:
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.api_key = api_key

def _request(self, method: str, path: str, body: dict | None) -> tuple[int, dict[str, Any]]:
url = f"{self.base_url}{path}"
data = json.dumps(body).encode("utf-8") if body is not None else None
headers = {"Content-Type": "application/json"}
if self.api_key:
headers["Authorization"] = f"Bearer {self.api_key}"
req = urllib.request.Request(
url,
data=data,
headers={"Content-Type": "application/json"},
headers=headers,
method=method,
)
try:
Expand Down Expand Up @@ -124,8 +145,16 @@ def _read_json(resp) -> dict[str, Any]:
return {"error": "BadResponse", "message": raw.decode("utf-8", "replace")}


def make_client(server_url: str | None, configs: list[str] | None = None) -> EngineClient:
"""Pick a backend: HTTP when ``server_url`` is given, else in-process."""
def make_client(
server_url: str | None,
configs: list[str] | None = None,
api_key: str | None = None,
) -> EngineClient:
"""Pick a backend: HTTP when ``server_url`` is given, else in-process.

``api_key`` 缺省读环境变量 ``AGENT_MEMORY_API_KEY``——让 key 不必出现在
shell history 与 ``ps`` 输出里。
"""
if server_url:
return HttpClient(server_url)
return HttpClient(server_url, api_key=api_key or os.environ.get("AGENT_MEMORY_API_KEY", ""))
return InProcessClient(configs)
143 changes: 143 additions & 0 deletions bootstrap/core/auth_middleware.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,143 @@
"""请求作用域的认证上下文——凭据提取 + ContextVar 建立/清理。

各 surface(HTTP / MCP / CLI 直连)用同一条中间件:把本形态的凭据材料归一成
:class:`~security.types.Credentials`,交给装配好的 ``Authenticator``,把产出的
``AuthContext`` 挂进 ContextVar 供 ``handler.dispatch`` 读取。

**本模块不决定认证策略**——模式(dev / trusted / api_key)由配置在装配期选定,
这里只负责「在正确的时机调用它、并保证退出时清理干净」。
"""

from __future__ import annotations

import os
import sys
from contextlib import contextmanager
from importlib import import_module
from typing import Any, Iterator, Mapping

_SRC = os.path.join(
os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "src"
)
if _SRC not in sys.path:
sys.path.append(_SRC)

_auth_module = import_module("common.type_def.auth")
reset_current = _auth_module.reset_current
set_current = _auth_module.set_current

Scope = import_module("common.type_def").Scope
AuditEvent = import_module("common.type_def").AuditEvent
_errors = import_module("common.errors")
AuthenticationError = _errors.AuthenticationError
RateLimitedError = _errors.RateLimitedError
Credentials = import_module("security.types").Credentials

_BEARER = "bearer "

_RATE_LIMITED = "too many requests"


def credentials_from_headers(headers: Mapping[str, Any], peer_address: str = "") -> Credentials:
"""从 HTTP header 提取凭据。

HTTP header 名大小写不敏感(RFC 9110 §5.1)。``http.client.HTTPMessage`` 的
``get`` 自己会做不敏感匹配,但传给 authenticator 的是普通 Mapping——故在这里
统一归一成小写键,authenticator 侧按小写常量查,两边不必各写一次 ``.lower()``。
"""
normalized = {str(k).lower(): str(v) for k, v in headers.items()}

api_key = ""
auth = normalized.get("authorization", "")
bearer_len = len(_BEARER)
if auth[:bearer_len].lower() == _BEARER:
api_key = auth[bearer_len:].strip()
if not api_key:
api_key = normalized.get("x-api-key", "").strip()

return Credentials(api_key=api_key, headers=normalized, peer_address=peer_address)


@contextmanager
def authenticated(
authenticator, credentials, audit=None, limiter=None, *, argon2_guard=None
) -> Iterator[Any]:
"""在请求作用域内建立可信认证上下文;退出时**必定** reset。

reset 放 ``finally`` 是硬性要求:``ThreadingHTTPServer`` 每请求一线程,
但线程可能被池化复用;漏 reset 会让下一个请求继承上一个请求的身份——
最严重的一类越权。

``authenticate`` 故意放在 ``try`` 之外:认证失败时没有 token 可 reset,
放进 try 会需要一个 ``token = None`` 的分支判断,反而更容易写错。

``limiter`` 在 ``authenticate`` **之前**执行(§8.1):认证本身就是要保护的
资源——API_KEY 模式下每次 authenticate 跑一次 Argon2id verify(128 MiB ×
time_cost=4),放在认证之后限流就等于「先让攻击者把 CPU 用掉,再告诉他
超限了」。``limiter=None`` 表示不限流(进程内直连 / MCP stdio 无网络对端)。

``argon2_guard`` 是进程级并发上限(审计 P1-3):IP 桶限请求速率,限不住
「同时在跑的 Argon2 verify 数」。耗尽即 429,在 limiter 之后、authenticate
之前执行。acquire 成功后用 ``finally`` 释放;``None`` 表示不限(DEV 模式或
调用方确信不跑 Argon2)。
"""
if limiter is not None and not limiter.allow(credentials.peer_address):
_record_denial(audit, authenticator, credentials, "rate_limit")
raise RateLimitedError(_RATE_LIMITED)

guard_acquired = False
if argon2_guard is not None:
if not argon2_guard.acquire():
_record_denial(audit, authenticator, credentials, "argon2_concurrency")
raise RateLimitedError(_RATE_LIMITED)
guard_acquired = True

try:
ctx = authenticator.authenticate(credentials)
except AuthenticationError:
_record_denial(audit, authenticator, credentials, "authenticate")
raise
finally:
if guard_acquired:
argon2_guard.release()

token = set_current(ctx)
try:
yield ctx
finally:
reset_current(token)


def _record_denial(audit, authenticator, credentials, action) -> None:
"""入口拒绝落一条审计(security.md §7.2):``action`` 区分限流与认证失败。

每次拒绝都记,无阈值聚合——限流器的计数器目前只用于准入判断,不对外暴露
统计;要做「同一 peer 连续失败 N 次告警」还需要一个独立的失败计数维度
(限流桶按请求数计,不区分成功与失败),那是可观测性设计,不在本期。

``actor`` 是空 ``Scope()``——身份未知,**不可用调用方声明的任何值填充**。
``detail`` 里不放 api_key、不放 key 前缀(§7.5 PII 脱敏),也不放桶余量
(那能用来反推限流参数)。

暂不记录认证失败的细分原因(``missing_credentials`` / ``unknown_principal`` /
``bad_gateway_key``):三个 authenticator 都刻意只抛同一个笼统消息,要拿到
细分原因得在 authenticator 侧另开一条只进审计的通道。那是独立设计,
不顺手塞进本期。
"""
if audit is None:
return
try:
audit.record(
AuditEvent(
actor=Scope(),
action=action,
decision="deny",
layer="security",
detail={
"mode": authenticator.mode().value,
"peer": credentials.peer_address,
},
)
)
except Exception: # pragma: no cover - 审计后端故障不该把 401/429 变成 500
pass
Loading