Coverage for agentos/enterprise/auth.py: 52%
206 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 12:20 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 12:20 +0800
1"""
2AgentOS Enterprise — SSO & RBAC.
4功能:
5 - RBAC 角色模型(admin / developer / viewer / agent)
6 - 权限定义与校验
7 - SSO 集成接口(OIDC / SAML 抽象)
8 - JWT Token 签发与验证
9 - 会话管理
10"""
12from __future__ import annotations
14import hashlib
15import hmac
16import json
17import time
18from dataclasses import dataclass, field
19from enum import StrEnum
21# ── 权限系统 ──
24class Permission(StrEnum):
25 """细粒度权限定义。"""
27 # Agent
28 AGENT_CREATE = "agent:create"
29 AGENT_READ = "agent:read"
30 AGENT_UPDATE = "agent:update"
31 AGENT_DELETE = "agent:delete"
32 AGENT_RUN = "agent:run"
33 # Tools
34 TOOLS_LIST = "tools:list"
35 TOOLS_EXECUTE = "tools:execute"
36 TOOLS_MANAGE = "tools:manage"
37 # API Keys
38 KEYS_CREATE = "keys:create"
39 KEYS_READ = "keys:read"
40 KEYS_REVOKE = "keys:revoke"
41 # Tenants
42 TENANT_READ = "tenant:read"
43 TENANT_MANAGE = "tenant:manage"
44 # Audit
45 AUDIT_READ = "audit:read"
46 AUDIT_EXPORT = "audit:export"
47 # Admin
48 ADMIN_ALL = "admin:*"
49 SYSTEM_CONFIG = "system:config"
52class Role(StrEnum):
53 """预定义角色。"""
55 ADMIN = "admin"
56 DEVELOPER = "developer"
57 VIEWER = "viewer"
58 AGENT = "agent"
61# 角色权限映射
62ROLE_PERMISSIONS: dict[Role, set[Permission]] = {
63 Role.ADMIN: set(Permission), # 全部权限
64 Role.DEVELOPER: {
65 Permission.AGENT_CREATE,
66 Permission.AGENT_READ,
67 Permission.AGENT_UPDATE,
68 Permission.AGENT_RUN,
69 Permission.TOOLS_LIST,
70 Permission.TOOLS_EXECUTE,
71 Permission.KEYS_CREATE,
72 Permission.KEYS_READ,
73 Permission.AUDIT_READ,
74 },
75 Role.VIEWER: {
76 Permission.AGENT_READ,
77 Permission.TOOLS_LIST,
78 Permission.KEYS_READ,
79 Permission.AUDIT_READ,
80 Permission.TENANT_READ,
81 },
82 Role.AGENT: {
83 Permission.AGENT_RUN,
84 Permission.TOOLS_EXECUTE,
85 },
86}
89@dataclass
90class User:
91 """用户实体。"""
93 user_id: str
94 username: str
95 email: str
96 roles: list[Role]
97 tenant_id: str
98 custom_permissions: set[Permission] = field(default_factory=set)
99 disabled: bool = False
100 created_at: float = field(default_factory=time.time)
101 metadata: dict = field(default_factory=dict)
104class RBACEngine:
105 """RBAC 权限引擎。
107 特性:
108 - 角色 + 自定义权限叠加
109 - 权限继承(admin 拥有全部)
110 - 批量权限检查
111 - 权限审计日志
112 """
114 def __init__(self):
115 self._custom_roles: dict[str, set[Permission]] = {}
117 def get_permissions(self, user: User) -> set[Permission]:
118 """获取用户的所有有效权限。"""
119 if user.disabled:
120 return set()
122 perms: set[Permission] = set(user.custom_permissions)
124 for role in user.roles:
125 perms |= ROLE_PERMISSIONS.get(role, set())
127 # Admin 自动获得全部
128 if Role.ADMIN in user.roles:
129 perms = set(Permission)
131 return perms
133 def check_permission(self, user: User, permission: Permission) -> bool:
134 """检查用户是否有某权限。"""
135 return permission in self.get_permissions(user)
137 def check_permissions(
138 self, user: User, permissions: list[Permission]
139 ) -> dict[Permission, bool]:
140 """批量权限检查。"""
141 user_perms = self.get_permissions(user)
142 return {p: p in user_perms for p in permissions}
144 def has_any(self, user: User, permissions: list[Permission]) -> bool:
145 """用户是否拥有任一权限。"""
146 user_perms = self.get_permissions(user)
147 return bool(user_perms & set(permissions))
149 def has_all(self, user: User, permissions: list[Permission]) -> bool:
150 """用户是否拥有全部权限。"""
151 user_perms = self.get_permissions(user)
152 return set(permissions).issubset(user_perms)
154 def register_custom_role(self, name: str, permissions: set[Permission]):
155 """注册自定义角色。"""
156 self._custom_roles[name] = permissions
158 def get_role_permissions(self, role: Role) -> set[Permission]:
159 return ROLE_PERMISSIONS.get(role, set())
162# ── SSO 集成 ──
165@dataclass
166class OIDCConfig:
167 """OIDC 提供商配置。"""
169 issuer: str # 如 "https://accounts.google.com"
170 client_id: str
171 client_secret: str
172 redirect_uri: str
173 scopes: list[str] = field(default_factory=lambda: ["openid", "email", "profile"])
174 authorization_endpoint: str = ""
175 token_endpoint: str = ""
176 userinfo_endpoint: str = ""
177 jwks_uri: str = ""
180@dataclass
181class SAMLConfig:
182 """SAML 提供商配置。"""
184 idp_entity_id: str
185 idp_sso_url: str
186 idp_certificate: str
187 sp_entity_id: str
188 sp_acs_url: str
189 name_id_format: str = "urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress"
192@dataclass
193class SSOUser:
194 """SSO 返回的用户信息。"""
196 external_id: str
197 email: str
198 display_name: str
199 provider: str # "oidc" / "saml"
200 raw_claims: dict = field(default_factory=dict)
203class SSOProvider:
204 """SSO 抽象层 — OIDC / SAML 统一接口。"""
206 @staticmethod
207 def build_oidc_login_url(config: OIDCConfig, state: str = "", nonce: str = "") -> str:
208 """构建 OIDC 登录 URL。"""
209 import urllib.parse
211 params = {
212 "response_type": "code",
213 "client_id": config.client_id,
214 "redirect_uri": config.redirect_uri,
215 "scope": " ".join(config.scopes),
216 "state": state or _rand_str(16),
217 "nonce": nonce or _rand_str(16),
218 }
219 ep = config.authorization_endpoint or f"{config.issuer.rstrip('/')}/authorize"
220 return f"{ep}?{urllib.parse.urlencode(params)}"
222 @staticmethod
223 def build_saml_login_url(config: SAMLConfig, relay_state: str = "") -> str:
224 """构建 SAML 登录 URL(SAMLRequest Base64)。"""
225 import base64
226 import uuid
228 saml_request = (
229 f'<?xml version="1.0" encoding="UTF-8"?>'
230 f'<samlp:AuthnRequest xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol"'
231 f' ID="_{uuid.uuid4().hex}" Version="2.0"'
232 f' IssueInstant="{time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())}"'
233 f' Destination="{config.idp_sso_url}"'
234 f' AssertionConsumerServiceURL="{config.sp_acs_url}">'
235 f'<saml:Issuer xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion">'
236 f"{config.sp_entity_id}</saml:Issuer>"
237 f"</samlp:AuthnRequest>"
238 )
239 encoded = base64.b64encode(saml_request.encode()).decode()
240 import urllib.parse
242 params = {"SAMLRequest": encoded}
243 if relay_state:
244 params["RelayState"] = relay_state
245 return f"{config.idp_sso_url}?{urllib.parse.urlencode(params)}"
247 @staticmethod
248 async def exchange_oidc_code(config: OIDCConfig, code: str) -> SSOUser | None:
249 """用 OIDC authorization_code 交换 token 并获取用户信息。(需要 httpx)"""
250 try:
251 import httpx
252 except ImportError:
253 return None
255 token_ep = config.token_endpoint or f"{config.issuer.rstrip('/')}/token"
256 async with httpx.AsyncClient() as client:
257 resp = await client.post(
258 token_ep,
259 data={
260 "grant_type": "authorization_code",
261 "code": code,
262 "redirect_uri": config.redirect_uri,
263 "client_id": config.client_id,
264 "client_secret": config.client_secret,
265 },
266 )
267 if resp.status_code != 200:
268 return None
269 token_data = resp.json()
270 access_token = token_data.get("access_token")
272 userinfo_ep = config.userinfo_endpoint or f"{config.issuer.rstrip('/')}/userinfo"
273 resp2 = await client.get(
274 userinfo_ep,
275 headers={
276 "Authorization": f"Bearer {access_token}",
277 },
278 )
279 if resp2.status_code != 200:
280 return None
281 info = resp2.json()
282 return SSOUser(
283 external_id=info.get("sub", ""),
284 email=info.get("email", ""),
285 display_name=info.get("name", info.get("preferred_username", "")),
286 provider="oidc",
287 raw_claims=info,
288 )
289 return None
292# ── 会话管理 ──
295@dataclass
296class Session:
297 """用户会话。"""
299 session_id: str
300 user_id: str
301 tenant_id: str
302 roles: list[Role]
303 created_at: float = field(default_factory=time.time)
304 expires_at: float = field(default_factory=lambda: time.time() + 3600) # 1 小时
305 ip_address: str = ""
306 user_agent: str = ""
308 def is_expired(self) -> bool:
309 return time.time() > self.expires_at
312class SessionStore:
313 """内存会话存储(生产环境应替换为 Redis)。"""
315 def __init__(self):
316 self._sessions: dict[str, Session] = {}
318 def create(self, user: User, ip: str = "", ua: str = "", ttl: int = 3600) -> Session:
319 import uuid
321 session = Session(
322 session_id=f"sess_{uuid.uuid4().hex[:16]}",
323 user_id=user.user_id,
324 tenant_id=user.tenant_id,
325 roles=user.roles,
326 expires_at=time.time() + ttl,
327 ip_address=ip,
328 user_agent=ua,
329 )
330 self._sessions[session.session_id] = session
331 return session
333 def get(self, session_id: str) -> Session | None:
334 s = self._sessions.get(session_id)
335 if s and s.is_expired():
336 del self._sessions[session_id]
337 return None
338 return s
340 def revoke(self, session_id: str):
341 self._sessions.pop(session_id, None)
343 def revoke_user_sessions(self, user_id: str):
344 to_remove = [sid for sid, s in self._sessions.items() if s.user_id == user_id]
345 for sid in to_remove:
346 del self._sessions[sid]
348 def stats(self) -> dict:
349 active = sum(1 for s in self._sessions.values() if not s.is_expired())
350 return {"total": len(self._sessions), "active": active}
353# ── JWT ──
356class JWTManager:
357 """简易 JWT 签发/验证(无外部依赖)。
359 生产环境建议使用 PyJWT / jwcrypto。
360 """
362 def __init__(self, secret: str):
363 self.secret = secret
365 def encode(self, payload: dict, ttl: int = 3600) -> str:
366 """签发 JWT。"""
367 import base64
369 header = {"alg": "HS256", "typ": "JWT"}
370 claims = {
371 **payload,
372 "iat": int(time.time()),
373 "exp": int(time.time()) + ttl,
374 }
375 segments = [
376 base64.urlsafe_b64encode(json.dumps(header).encode()).rstrip(b"=").decode(),
377 base64.urlsafe_b64encode(json.dumps(claims).encode()).rstrip(b"=").decode(),
378 ]
379 signing_input = ".".join(segments)
380 sig = hmac.new(self.secret.encode(), signing_input.encode(), hashlib.sha256).digest()
381 segments.append(base64.urlsafe_b64encode(sig).rstrip(b"=").decode())
382 return ".".join(segments)
384 def decode(self, token: str) -> dict | None:
385 """验证并解码 JWT。"""
386 import base64
388 try:
389 parts = token.split(".")
390 if len(parts) != 3:
391 return None
393 header_b64, payload_b64, sig_b64 = parts
394 signing_input = f"{header_b64}.{payload_b64}"
396 # Verify signature
397 expected_sig = (
398 base64.urlsafe_b64encode(
399 hmac.new(self.secret.encode(), signing_input.encode(), hashlib.sha256).digest()
400 )
401 .rstrip(b"=")
402 .decode()
403 )
405 if not hmac.compare_digest(sig_b64, expected_sig):
406 return None
408 # Decode payload
409 payload = json.loads(base64.urlsafe_b64decode(payload_b64 + "==").decode())
411 # Check expiration
412 if payload.get("exp", 0) < time.time():
413 return None
415 return payload
416 except Exception:
417 return None
420# ── 工具函数 ──
423def _rand_str(n: int) -> str:
424 import secrets
426 return secrets.token_hex(n // 2 + 1)[:n]
429def require_permission(permission: Permission):
430 """装饰器:要求调用者拥有指定权限。(示例用途)"""
432 def decorator(func):
433 def wrapper(*args, **kwargs):
434 # 实际使用时会从上下文获取当前用户
435 raise NotImplementedError("权限检查需在框架中间件中实现")
437 return wrapper
439 return decorator