""" Shared pytest fixtures for integration tests that exercise the live API service. """ from __future__ import annotations import asyncio import json import os import subprocess import sys import uuid from collections.abc import AsyncGenerator from pathlib import Path import anyio import httpx import pytest import pytest_asyncio from dotenv import load_dotenv PROJECT_ROOT = Path(__file__).resolve().parents[2] if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) load_dotenv(PROJECT_ROOT / ".env", override=False) load_dotenv(PROJECT_ROOT / "test/.env.test", override=False) API_BASE_URL = os.getenv("TEST_BASE_URL", "http://localhost:5050").rstrip("/") ADMIN_LOGIN = os.getenv("TEST_USERNAME") ADMIN_PASSWORD = os.getenv("TEST_PASSWORD") _ADMIN_TOKEN_CACHE: str | None = None # 防止并发首次登录造成雪崩(触发 429 限流 / 423 账户锁定) _login_lock = asyncio.Lock() HTTP_TIMEOUT = httpx.Timeout(60.0, connect=5.0) SANDBOX_CONTAINER_PREFIX = os.getenv("YUXI_SANDBOX_CONTAINER_PREFIX", "yuxi-sandbox") @pytest.fixture(scope="session", autouse=True) def ensure_live_api_schema(): if not ADMIN_LOGIN or not ADMIN_PASSWORD: return async def run_schema_setup() -> None: from yuxi.storage.postgres.manager import pg_manager pg_manager.initialize() await pg_manager.create_tables() await pg_manager.ensure_business_schema() await pg_manager.ensure_knowledge_schema() await pg_manager.ensure_external_schema() await pg_manager.ensure_scheduler_schema() await pg_manager.ensure_channel_schema() anyio.run(run_schema_setup) def _require_admin_credentials() -> tuple[str, str]: if not ADMIN_LOGIN or not ADMIN_PASSWORD: pytest.skip("Integration credentials are not configured via TEST_USERNAME / TEST_PASSWORD.") return ADMIN_LOGIN, ADMIN_PASSWORD @pytest_asyncio.fixture(scope="function") async def test_client() -> AsyncGenerator[httpx.AsyncClient, None]: async with httpx.AsyncClient(base_url=API_BASE_URL, timeout=HTTP_TIMEOUT, follow_redirects=True) as client: yield client @pytest_asyncio.fixture(scope="function") async def admin_token() -> str: """获取管理员 token,带缓存预检与限流退避重试。 设计要点: - 缓存命中时先用 GET /api/auth/me 校验 token 有效性(GET 不触发登录限流), 避免每个用例都重新登录造成雪崩。 - token 失效才重新登录;登录遇到 429 按 Retry-After 退避重试,最多 2 次。 - 模块级 _login_lock 防止多用例并发首次登录。 """ global _ADMIN_TOKEN_CACHE async with httpx.AsyncClient( base_url=API_BASE_URL, timeout=HTTP_TIMEOUT, follow_redirects=True, ) as client: # [1] 预检:缓存命中时用 GET /api/auth/me 校验 token 是否仍然有效 if _ADMIN_TOKEN_CACHE: try: me_resp = await client.get( "/api/auth/me", headers={"Authorization": f"Bearer {_ADMIN_TOKEN_CACHE}"}, ) except httpx.HTTPError as exc: pytest.fail(f"Failed to verify admin token via /api/auth/me: {exc}") if me_resp.status_code == 200: return _ADMIN_TOKEN_CACHE # type: ignore[return-value] if me_resp.status_code == 423: pytest.fail( "Admin account is login-locked. Unlock it in DB (reset login_failed_count / " "login_locked_until) and restart api-dev before re-running tests." ) # 401 / 其他状态:token 已失效,清缓存走重新登录流程 _ADMIN_TOKEN_CACHE = None # [2] 登录:缓存失效或首次。加锁防止并发首次登录雪崩 username, password = _require_admin_credentials() async with _login_lock: # 双检:持锁后再次确认缓存(其他用例可能已登录成功) if _ADMIN_TOKEN_CACHE: return _ADMIN_TOKEN_CACHE # type: ignore[return-value] response = await _do_login_with_backoff(client, username, password) if response.status_code == 401: first_run_response = await client.get("/api/auth/check-first-run") if first_run_response.status_code == 200 and first_run_response.json().get("first_run", False): pytest.fail( "Super admin account has not been initialized. Complete `/api/auth/initialize` before " "running integration tests." ) if response.status_code != 200: pytest.fail(f"Failed to authenticate as admin (status={response.status_code}): {response.text}") token = response.json().get("access_token") if not token: pytest.fail("Admin authentication did not return an access token.") _ADMIN_TOKEN_CACHE = token return token async def _do_login_with_backoff( client: httpx.AsyncClient, username: str, password: str, ) -> httpx.Response: """执行登录,遇到 429 按 Retry-After 退避重试,最多 2 次。""" max_retries = 2 response = await client.post( "/api/auth/token", data={"username": username, "password": password}, ) for attempt in range(max_retries): if response.status_code != 429: break # 解析 Retry-After(服务端返回的 details.retry_after 优先,header 次之) retry_after: float = 5.0 try: body = response.json() retry_after = float(body.get("error", {}).get("details", {}).get("retry_after", retry_after)) except (ValueError, KeyError): header_val = response.headers.get("Retry-After") if header_val: try: retry_after = float(header_val) except ValueError: pass await anyio.sleep(retry_after + 1.0) response = await client.post( "/api/auth/token", data={"username": username, "password": password}, ) return response @pytest.fixture(scope="function") def admin_headers(admin_token: str) -> dict[str, str]: return {"Authorization": f"Bearer {admin_token}"} @pytest.fixture(scope="session", autouse=True) def cleanup_test_knowledge_databases(): async def run_cleanup() -> None: global _ADMIN_TOKEN_CACHE if not ADMIN_LOGIN or not ADMIN_PASSWORD: return if not _ADMIN_TOKEN_CACHE: async with httpx.AsyncClient( base_url=API_BASE_URL, timeout=HTTP_TIMEOUT, follow_redirects=True, ) as bootstrap_client: response = await bootstrap_client.post( "/api/auth/token", data={"username": ADMIN_LOGIN, "password": ADMIN_PASSWORD}, ) if response.status_code != 200: return token = response.json().get("access_token") if not token: return _ADMIN_TOKEN_CACHE = token headers = {"Authorization": f"Bearer {_ADMIN_TOKEN_CACHE}"} async with httpx.AsyncClient(base_url=API_BASE_URL, timeout=HTTP_TIMEOUT, follow_redirects=True) as client: try: list_response = await client.get("/api/knowledge/databases", headers=headers) except Exception as exc: print(f"Warning: Failed to list knowledge databases for cleanup: {exc}") return if list_response.status_code != 200: return databases = list_response.json().get("databases", []) prefixes = ("pytest_", "py_test") for entry in databases: name = entry.get("name") or "" slug = entry.get("slug") if not slug or not isinstance(name, str) or not name.startswith(prefixes): continue try: delete_response = await client.delete(f"/api/knowledge/databases/{slug}", headers=headers) if delete_response.status_code not in (200, 404): print(f"Warning: Failed to cleanup knowledge database {slug}: {delete_response.text}") except Exception as exc: print(f"Warning: Exception during cleanup of {slug}: {exc}") try: anyio.run(run_cleanup) except Exception as exc: print(f"Warning: Exception during session cleanup startup: {exc}") yield try: anyio.run(run_cleanup) except Exception as exc: print(f"Warning: Exception during session cleanup teardown: {exc}") def _docker_api_request(method: str, path: str) -> list[dict] | dict: cmd = [ "curl", "-sS", "--unix-socket", os.getenv("YUXI_DOCKER_API_SOCKET", "/var/run/docker.sock"), "-X", method, f"{os.getenv('YUXI_DOCKER_API_BASE', 'http://localhost').rstrip('/')}{path}", ] result = subprocess.run(cmd, capture_output=True, text=True, timeout=15, check=False) if result.returncode != 0: raise RuntimeError(result.stderr.strip() or "Docker API request failed") text = (result.stdout or "").strip() if not text: return {} return json.loads(text) def _cleanup_sandbox_containers() -> None: try: containers = _docker_api_request("GET", "/containers/json?all=true") except Exception as exc: print(f"Warning: Failed to list sandbox containers for cleanup: {exc}") return for container in containers if isinstance(containers, list) else []: names = container.get("Names") or [] if not any(name.lstrip("/").startswith(f"{SANDBOX_CONTAINER_PREFIX}-") for name in names): continue container_id = container.get("Id") if not container_id: continue try: _docker_api_request("POST", f"/containers/{container_id}/stop?t=2") except Exception: pass try: _docker_api_request("DELETE", f"/containers/{container_id}?force=true") except Exception as exc: print(f"Warning: Failed to cleanup sandbox container {container_id[:12]}: {exc}") @pytest.fixture(scope="session", autouse=True) def cleanup_test_sandboxes(): _cleanup_sandbox_containers() yield _cleanup_sandbox_containers() @pytest_asyncio.fixture(scope="function") async def standard_user(test_client: httpx.AsyncClient, admin_headers: dict[str, str]) -> AsyncGenerator[dict, None]: username = f"pytest_user_{uuid.uuid4().hex[:8]}" password = f"Pw!{uuid.uuid4().hex[:8]}" # 用户隔离重构后所有登录用户必须绑定部门,创建时显式指定一个已存在部门 dept_response = await test_client.get("/api/departments", headers=admin_headers) if dept_response.status_code != 200 or not dept_response.json(): pytest.fail(f"No department available to bind standard user: {dept_response.text}") department_id = dept_response.json()[0]["id"] response = await test_client.post( "/api/auth/users", json={"username": username, "password": password, "role": "user", "department_id": department_id}, headers=admin_headers, ) if response.status_code != 200: pytest.fail(f"Failed to create standard user (status={response.status_code}): {response.text}") user_payload = response.json() login_response = await test_client.post( "/api/auth/token", data={"username": user_payload["uid"], "password": password}, ) if login_response.status_code != 200: pytest.fail( f"Failed to authenticate as standard user (status={login_response.status_code}): {login_response.text}" ) access_token = login_response.json().get("access_token") if not access_token: pytest.fail("Standard user login succeeded but no access token was returned.") try: yield { "user": user_payload, "password": password, "headers": {"Authorization": f"Bearer {access_token}"}, } finally: cleanup_error = None for _ in range(3): response = await test_client.delete(f"/api/auth/users/{user_payload['id']}", headers=admin_headers) if response.status_code in (200, 404): cleanup_error = None break cleanup_error = response await anyio.sleep(0.3) if cleanup_error is not None: assert cleanup_error.status_code == 200, ( f"Failed to cleanup test user {user_payload['uid']}: {cleanup_error.text}" ) @pytest_asyncio.fixture(scope="function") async def knowledge_database( test_client: httpx.AsyncClient, admin_headers: dict[str, str], ) -> AsyncGenerator[dict, None]: import time unique_id = uuid.uuid4().hex timestamp = int(time.time() * 1000000) db_name = f"pytest_kb_{timestamp}_{unique_id}" kb_id = None try: create_response = await test_client.post( "/api/knowledge/databases", json={ "database_name": db_name, "description": "Pytest managed knowledge base", "embedding_model_spec": "siliconflow-cn:Pro/BAAI/bge-m3", "kb_type": "milvus", "additional_params": {}, }, headers=admin_headers, ) if create_response.status_code == 200: db_payload = create_response.json() kb_id = db_payload["kb_id"] elif create_response.status_code == 409: error_detail = create_response.json().get("detail", "") pytest.fail(f"Knowledge database name conflict: {error_detail}. Please clean up old test databases first.") else: pytest.fail( f"Failed to create knowledge database (status={create_response.status_code}): {create_response.text}" ) yield db_payload finally: if kb_id: try: delete_response = await test_client.delete(f"/api/knowledge/databases/{kb_id}", headers=admin_headers) if delete_response.status_code != 200: print(f"Warning: Failed to cleanup knowledge database {kb_id}: {delete_response.text}") except Exception as exc: print(f"Warning: Exception during cleanup of {kb_id}: {exc}")