Token lifecycle manager with AT auto-refresh
| 13 | |
| 14 | |
| 15 | class TokenManager: |
| 16 | """Token lifecycle manager with AT auto-refresh""" |
| 17 | |
| 18 | def __init__(self, db: Database, flow_client: FlowClient): |
| 19 | self.db = db |
| 20 | self.flow_client = flow_client |
| 21 | self._refresh_lock_guard = asyncio.Lock() |
| 22 | self._project_lock_guard = asyncio.Lock() |
| 23 | self._refresh_locks: dict[int, asyncio.Lock] = {} |
| 24 | self._project_locks: dict[int, asyncio.Lock] = {} |
| 25 | self._refresh_futures: dict[int, asyncio.Task] = {} |
| 26 | self._at_validation_cache: dict[int, float] = {} |
| 27 | self._protocol_refresher_task: Optional[asyncio.Task] = None |
| 28 | |
| 29 | async def _get_token_lock( |
| 30 | self, |
| 31 | lock_map: dict[int, asyncio.Lock], |
| 32 | guard: asyncio.Lock, |
| 33 | token_id: int, |
| 34 | ) -> asyncio.Lock: |
| 35 | """按 token 维度获取锁,避免不同 token 之间串行阻塞。""" |
| 36 | async with guard: |
| 37 | lock = lock_map.get(token_id) |
| 38 | if lock is None: |
| 39 | lock = asyncio.Lock() |
| 40 | lock_map[token_id] = lock |
| 41 | return lock |
| 42 | |
| 43 | def _get_project_pool_size(self) -> int: |
| 44 | """读取当前生效的单 Token 项目池大小配置。""" |
| 45 | try: |
| 46 | return max(1, min(50, int(config.personal_project_pool_size or 4))) |
| 47 | except Exception: |
| 48 | return 4 |
| 49 | |
| 50 | def _sort_projects(self, projects: List[Project]) -> List[Project]: |
| 51 | """Sort projects in a stable order for round-robin selection.""" |
| 52 | return sorted(projects, key=lambda project: (project.id or 0, project.project_id)) |
| 53 | |
| 54 | def _normalize_project_name_base(self, project_name: Optional[str] = None) -> str: |
| 55 | """Normalize a project base name for pooled creation.""" |
| 56 | raw_name = (project_name or "").strip() |
| 57 | if raw_name: |
| 58 | parts = raw_name.rsplit(" ", 1) |
| 59 | if len(parts) == 2 and parts[1].startswith("P") and parts[1][1:].isdigit(): |
| 60 | return parts[0] |
| 61 | return raw_name |
| 62 | return datetime.now().strftime("%b %d - %H:%M") |
| 63 | |
| 64 | def _build_project_name(self, pool_index: int, base_name: Optional[str] = None) -> str: |
| 65 | """Build a project name for the pool.""" |
| 66 | normalized_base = self._normalize_project_name_base(base_name) |
| 67 | return f"{normalized_base} P{pool_index}" |
| 68 | |
| 69 | def _normalize_protocol_mode(self, value: Optional[str]) -> str: |
| 70 | mode = (value or "session").strip().lower() |
| 71 | return "protocol" if mode == "protocol" else "session" |
| 72 |