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

1""" 

2AgentOS Enterprise — SSO & RBAC. 

3 

4功能: 

5 - RBAC 角色模型(admin / developer / viewer / agent) 

6 - 权限定义与校验 

7 - SSO 集成接口(OIDC / SAML 抽象) 

8 - JWT Token 签发与验证 

9 - 会话管理 

10""" 

11 

12from __future__ import annotations 

13 

14import hashlib 

15import hmac 

16import json 

17import time 

18from dataclasses import dataclass, field 

19from enum import StrEnum 

20 

21# ── 权限系统 ── 

22 

23 

24class Permission(StrEnum): 

25 """细粒度权限定义。""" 

26 

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" 

50 

51 

52class Role(StrEnum): 

53 """预定义角色。""" 

54 

55 ADMIN = "admin" 

56 DEVELOPER = "developer" 

57 VIEWER = "viewer" 

58 AGENT = "agent" 

59 

60 

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} 

87 

88 

89@dataclass 

90class User: 

91 """用户实体。""" 

92 

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) 

102 

103 

104class RBACEngine: 

105 """RBAC 权限引擎。 

106 

107 特性: 

108 - 角色 + 自定义权限叠加 

109 - 权限继承(admin 拥有全部) 

110 - 批量权限检查 

111 - 权限审计日志 

112 """ 

113 

114 def __init__(self): 

115 self._custom_roles: dict[str, set[Permission]] = {} 

116 

117 def get_permissions(self, user: User) -> set[Permission]: 

118 """获取用户的所有有效权限。""" 

119 if user.disabled: 

120 return set() 

121 

122 perms: set[Permission] = set(user.custom_permissions) 

123 

124 for role in user.roles: 

125 perms |= ROLE_PERMISSIONS.get(role, set()) 

126 

127 # Admin 自动获得全部 

128 if Role.ADMIN in user.roles: 

129 perms = set(Permission) 

130 

131 return perms 

132 

133 def check_permission(self, user: User, permission: Permission) -> bool: 

134 """检查用户是否有某权限。""" 

135 return permission in self.get_permissions(user) 

136 

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} 

143 

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)) 

148 

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) 

153 

154 def register_custom_role(self, name: str, permissions: set[Permission]): 

155 """注册自定义角色。""" 

156 self._custom_roles[name] = permissions 

157 

158 def get_role_permissions(self, role: Role) -> set[Permission]: 

159 return ROLE_PERMISSIONS.get(role, set()) 

160 

161 

162# ── SSO 集成 ── 

163 

164 

165@dataclass 

166class OIDCConfig: 

167 """OIDC 提供商配置。""" 

168 

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 = "" 

178 

179 

180@dataclass 

181class SAMLConfig: 

182 """SAML 提供商配置。""" 

183 

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" 

190 

191 

192@dataclass 

193class SSOUser: 

194 """SSO 返回的用户信息。""" 

195 

196 external_id: str 

197 email: str 

198 display_name: str 

199 provider: str # "oidc" / "saml" 

200 raw_claims: dict = field(default_factory=dict) 

201 

202 

203class SSOProvider: 

204 """SSO 抽象层 — OIDC / SAML 统一接口。""" 

205 

206 @staticmethod 

207 def build_oidc_login_url(config: OIDCConfig, state: str = "", nonce: str = "") -> str: 

208 """构建 OIDC 登录 URL。""" 

209 import urllib.parse 

210 

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)}" 

221 

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 

227 

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 

241 

242 params = {"SAMLRequest": encoded} 

243 if relay_state: 

244 params["RelayState"] = relay_state 

245 return f"{config.idp_sso_url}?{urllib.parse.urlencode(params)}" 

246 

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 

254 

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") 

271 

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 

290 

291 

292# ── 会话管理 ── 

293 

294 

295@dataclass 

296class Session: 

297 """用户会话。""" 

298 

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 = "" 

307 

308 def is_expired(self) -> bool: 

309 return time.time() > self.expires_at 

310 

311 

312class SessionStore: 

313 """内存会话存储(生产环境应替换为 Redis)。""" 

314 

315 def __init__(self): 

316 self._sessions: dict[str, Session] = {} 

317 

318 def create(self, user: User, ip: str = "", ua: str = "", ttl: int = 3600) -> Session: 

319 import uuid 

320 

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 

332 

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 

339 

340 def revoke(self, session_id: str): 

341 self._sessions.pop(session_id, None) 

342 

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] 

347 

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} 

351 

352 

353# ── JWT ── 

354 

355 

356class JWTManager: 

357 """简易 JWT 签发/验证(无外部依赖)。 

358 

359 生产环境建议使用 PyJWT / jwcrypto。 

360 """ 

361 

362 def __init__(self, secret: str): 

363 self.secret = secret 

364 

365 def encode(self, payload: dict, ttl: int = 3600) -> str: 

366 """签发 JWT。""" 

367 import base64 

368 

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) 

383 

384 def decode(self, token: str) -> dict | None: 

385 """验证并解码 JWT。""" 

386 import base64 

387 

388 try: 

389 parts = token.split(".") 

390 if len(parts) != 3: 

391 return None 

392 

393 header_b64, payload_b64, sig_b64 = parts 

394 signing_input = f"{header_b64}.{payload_b64}" 

395 

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 ) 

404 

405 if not hmac.compare_digest(sig_b64, expected_sig): 

406 return None 

407 

408 # Decode payload 

409 payload = json.loads(base64.urlsafe_b64decode(payload_b64 + "==").decode()) 

410 

411 # Check expiration 

412 if payload.get("exp", 0) < time.time(): 

413 return None 

414 

415 return payload 

416 except Exception: 

417 return None 

418 

419 

420# ── 工具函数 ── 

421 

422 

423def _rand_str(n: int) -> str: 

424 import secrets 

425 

426 return secrets.token_hex(n // 2 + 1)[:n] 

427 

428 

429def require_permission(permission: Permission): 

430 """装饰器:要求调用者拥有指定权限。(示例用途)""" 

431 

432 def decorator(func): 

433 def wrapper(*args, **kwargs): 

434 # 实际使用时会从上下文获取当前用户 

435 raise NotImplementedError("权限检查需在框架中间件中实现") 

436 

437 return wrapper 

438 

439 return decorator