"""回归测试: /api/v1/predict 限流 + 短 session 模式。 验证: 1. 限流: 同 IP 超过 10 次/分钟返回 429 2. 限流: 不同 IP 独立计数 3. 限流: 滑动窗口过期后恢复 4. 短 session: predict 路由不持有 DB 连接 during LLM call """ from __future__ import annotations import asyncio import time import pytest from src.api.deps import _RateLimiter, rate_limit_predict class TestRateLimiter: """_RateLimiter 滑动窗口限流。""" def test_allows_within_limit(self): limiter = _RateLimiter(max_requests=10, window_seconds=60) for _ in range(10): assert limiter.is_allowed("192.168.1.1") def test_blocks_over_limit(self): limiter = _RateLimiter(max_requests=3, window_seconds=60) assert limiter.is_allowed("10.0.0.1") # 1 assert limiter.is_allowed("10.0.0.1") # 2 assert limiter.is_allowed("10.0.0.1") # 3 assert not limiter.is_allowed("10.0.0.1") # 4 → blocked def test_different_keys_independent(self): """不同 IP 的限流计数独立。""" limiter = _RateLimiter(max_requests=2, window_seconds=60) assert limiter.is_allowed("10.0.0.1") assert limiter.is_allowed("10.0.0.1") assert not limiter.is_allowed("10.0.0.1") # blocked # 不同 IP 仍允许 assert limiter.is_allowed("10.0.0.2") assert limiter.is_allowed("10.0.0.2") assert not limiter.is_allowed("10.0.0.2") # blocked def test_sliding_window_expires(self): """滑动窗口:过期后恢复。""" limiter = _RateLimiter(max_requests=2, window_seconds=1) assert limiter.is_allowed("10.0.0.1") assert limiter.is_allowed("10.0.0.1") assert not limiter.is_allowed("10.0.0.1") # blocked # 等待窗口过期 time.sleep(1.1) assert limiter.is_allowed("10.0.0.1") # 窗口过期,恢复 def test_cleans_expired_entries(self): """验证过期条目被清理(不会无限增长)。""" limiter = _RateLimiter(max_requests=100, window_seconds=1) for _ in range(50): limiter.is_allowed("10.0.0.1") # 验证内部状态 assert len(limiter._hits.get("10.0.0.1", [])) == 50 time.sleep(1.1) # 触发清理 limiter.is_allowed("10.0.0.1") # 过期条目应被清除,只剩新加入的 1 条 assert len(limiter._hits.get("10.0.0.1", [])) == 1 class TestShortReadSession: """short_read 上下文管理器。""" @pytest.mark.asyncio async def test_short_read_context_manager(self): """short_read 应作为 async context manager 工作。""" from src.db.base import short_read import inspect # 验证是 async context manager (通过 inspect 检查) assert inspect.isasyncgenfunction(short_read) or hasattr(short_read, "__wrapped__") # 验证可以调用并返回 context manager ctx = short_read() assert hasattr(ctx, "__aenter__") assert hasattr(ctx, "__aexit__") class TestDepsImports: """验证新依赖可正确导入。""" def test_rate_limit_predict_importable(self): from src.api.deps import rate_limit_predict assert callable(rate_limit_predict) def test_rate_limiter_importable(self): from src.api.deps import _RateLimiter, _predict_limiter assert isinstance(_predict_limiter, _RateLimiter) assert _predict_limiter.max_requests == 10 assert _predict_limiter.window_seconds == 60