test:删除测试代码 #1
@@ -1,67 +0,0 @@
|
||||
"""认证 E2E:登录取 token → /auth/me 依赖注入回查 sys_user → 反例(错密码/坏 token/已删用户)。
|
||||
|
||||
全程跑在同一个事件循环(httpx.ASGITransport),避免单例 async engine 跨循环。
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
import httpx
|
||||
|
||||
import main as appmod
|
||||
from config.database.mysql import get_engine, get_session_factory
|
||||
from model.sys_user import SysUser
|
||||
from repositories.sys_user import SysUserRepo
|
||||
from service.auth import hash_password
|
||||
|
||||
USER = "auth_test_u"
|
||||
PWD = "Test@1234"
|
||||
|
||||
|
||||
async def seed() -> int:
|
||||
async with get_session_factory()() as s:
|
||||
repo = SysUserRepo(s)
|
||||
existing = await repo.get_by_username(USER)
|
||||
if existing:
|
||||
return existing.id
|
||||
u = SysUser(username=USER, password_hash=hash_password(PWD), phone="13800001111",
|
||||
user_type="EMPLOYEE", employee_role="理财顾问", status="正常")
|
||||
return (await repo.add(u)).id
|
||||
|
||||
|
||||
async def cleanup(uid: int) -> None:
|
||||
async with get_session_factory()() as s:
|
||||
await SysUserRepo(s).delete(uid)
|
||||
|
||||
|
||||
async def run():
|
||||
uid = await seed()
|
||||
try:
|
||||
transport = httpx.ASGITransport(app=appmod.app)
|
||||
async with httpx.AsyncClient(transport=transport, base_url="http://test") as c:
|
||||
# 错密码 → 400
|
||||
r = await c.post("/api/auth/login", json={"username": USER, "password": "wrong"})
|
||||
assert r.status_code == 400 and r.json()["code"] == 400, r.text
|
||||
# 正常登录 → token
|
||||
r = await c.post("/api/auth/login", json={"username": USER, "password": PWD})
|
||||
assert r.status_code == 200, r.text
|
||||
token = r.json()["data"]["token"]
|
||||
assert token
|
||||
# /auth/me:依赖注入按 id 回查 sys_user
|
||||
r = await c.get("/api/auth/me", headers={"Authorization": f"Bearer {token}"})
|
||||
me = r.json()["data"]["user"]
|
||||
assert r.status_code == 200 and me["id"] == uid and me["employee_role"] == "理财顾问", r.text
|
||||
# 缺 token / 坏 token → 401
|
||||
assert (await c.get("/api/auth/me")).status_code == 401
|
||||
assert (await c.get("/api/auth/me", headers={"Authorization": "Bearer abc.def.ghi"})).status_code == 401
|
||||
# 删除用户后 token 回查失败 → 401
|
||||
await cleanup(uid)
|
||||
assert (await c.get("/api/auth/me", headers={"Authorization": f"Bearer {token}"})).status_code == 401
|
||||
# 不存在用户登录 → 400
|
||||
r = await c.post("/api/auth/login", json={"username": "no_such_user", "password": PWD})
|
||||
assert r.status_code == 400
|
||||
print("AUTH TESTS PASSED")
|
||||
finally:
|
||||
await get_engine().dispose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(run())
|
||||
-118
@@ -1,118 +0,0 @@
|
||||
"""tool.llm 自检:本地 Mock OpenAI 兼容服务,验证双模式 URL、重试、备用模型降级、embedding。
|
||||
|
||||
运行:python test_llm.py(纯 assert,无框架;可通过 exitcode 判断)
|
||||
"""
|
||||
import asyncio
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
|
||||
from config.settings import settings
|
||||
from tool.llm import LLMFailError, LLMClient
|
||||
|
||||
INDEX = json.dumps({"index": 0}).encode()
|
||||
|
||||
|
||||
class MockLLMServer(http.server.BaseHTTPRequestHandler):
|
||||
"""OpenAI 兼容 Mock:chat/embeddings/tags。model=boom:1b 返回 500 用于测降级。"""
|
||||
|
||||
def log_message(self, *args):
|
||||
pass
|
||||
|
||||
def _reply(self, code: int, obj: dict):
|
||||
body = json.dumps(obj).encode()
|
||||
self.send_response(code)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def do_GET(self): # ollama /api/tags 健康探测
|
||||
self._reply(200, {"models": [{"name": "qwen2.5:7b"}]})
|
||||
|
||||
def do_POST(self):
|
||||
n = int(self.headers.get("Content-Length", 0))
|
||||
req = json.loads(self.rfile.read(n))
|
||||
if self.path == "/v1/chat/completions":
|
||||
model = req["model"]
|
||||
if model.startswith("boom"):
|
||||
self._reply(500, {"error": "model overloaded"})
|
||||
return
|
||||
self._reply(200, {"choices": [{"message": {"content": f"reply:{model}"}}]})
|
||||
elif self.path == "/v1/embeddings":
|
||||
texts = req["input"]
|
||||
self._reply(200, {"data": [{"embedding": [float(i + 1)] * 4} for i in range(len(texts))]})
|
||||
else:
|
||||
self._reply(404, {"error": "not found"})
|
||||
|
||||
|
||||
def start_server() -> tuple[threading.Thread, int]:
|
||||
srv = http.server.ThreadingHTTPServer(("127.0.0.1", 0), MockLLMServer)
|
||||
port = srv.server_address[1]
|
||||
t = threading.Thread(target=srv.serve_forever, daemon=True)
|
||||
t.start()
|
||||
return t, port
|
||||
|
||||
|
||||
def cfg(base_url: str, **kw):
|
||||
"""从 .env 加载的 settings.llm 派生测试配置(只覆盖用例所需字段)。"""
|
||||
kw.setdefault("mode", "ollama")
|
||||
kw.setdefault("ollama_base", base_url)
|
||||
kw.setdefault("max_retries", 2)
|
||||
kw.setdefault("retry_backoff_sec", 0)
|
||||
return settings.llm.model_copy(update=kw)
|
||||
|
||||
|
||||
async def achat(c: LLMClient, text: str = "hi") -> str:
|
||||
return await c.chat([{"role": "user", "content": text}])
|
||||
|
||||
|
||||
def main():
|
||||
_, port = start_server()
|
||||
base = f"http://127.0.0.1:{port}"
|
||||
|
||||
# 1. ollama 模式 URL + 正常 chat
|
||||
c = LLMClient(cfg(base, ollama_chat_model="ok:1b"))
|
||||
assert c.is_ollama and c.base_url == f"{base}/v1", c.base_url
|
||||
out = asyncio.run(achat(c))
|
||||
assert out == "reply:ok:1b", out
|
||||
|
||||
# 2. 备用模型降级:主模型 500 → 自动切 fallback_chat_model
|
||||
c2 = LLMClient(cfg(base, ollama_chat_model="boom:1b", ollama_embed_model="bge-m3",
|
||||
fallback_chat_model="ok:1b"))
|
||||
out = asyncio.run(achat(c2))
|
||||
assert out == "reply:ok:1b", out
|
||||
|
||||
# 3. 全部失败 → LLMFailError(重试耗尽)
|
||||
c3 = LLMClient(cfg(base, ollama_chat_model="boom:1b", max_retries=1))
|
||||
try:
|
||||
asyncio.run(achat(c3))
|
||||
raise AssertionError("should have raised LLMFailError")
|
||||
except LLMFailError:
|
||||
pass
|
||||
|
||||
# 4. API 模式:端点/鉴权头/模型选择
|
||||
c4 = LLMClient(settings.llm.model_copy(update={
|
||||
"mode": "api", "api_base": base + "/v1", "api_key": "sk-test",
|
||||
"api_chat_model": "gpt-x", "max_retries": 1,
|
||||
}))
|
||||
assert not c4.is_ollama and c4.base_url == f"{base}/v1"
|
||||
assert c4._headers()["Authorization"] == "Bearer sk-test"
|
||||
assert c4.chat_model == "gpt-x"
|
||||
out = asyncio.run(achat(c4))
|
||||
assert out == "reply:gpt-x", out
|
||||
|
||||
# 5. embedding 批量
|
||||
c5 = LLMClient(cfg(base))
|
||||
vecs = asyncio.run(c5.embed(["基金", "债券"]))
|
||||
assert len(vecs) == 2 and len(vecs[0]) == 4, vecs
|
||||
assert asyncio.run(c5.embed_one("货币")) == [1.0] * 4
|
||||
|
||||
# 6. 健康探测(ollama /api/tags)
|
||||
asyncio.run(c5.check_health())
|
||||
|
||||
print("ALL LLM TESTS PASSED")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,11 +0,0 @@
|
||||
# Test your FastAPI endpoints
|
||||
|
||||
GET http://127.0.0.1:8000/
|
||||
Accept: application/json
|
||||
|
||||
###
|
||||
|
||||
GET http://127.0.0.1:8000/hello/User
|
||||
Accept: application/json
|
||||
|
||||
###
|
||||
@@ -1,44 +0,0 @@
|
||||
"""数据访问层自检:MySQL 可达且表已建则跑真实 CRUD;否则 SKIP(exit 0)。"""
|
||||
import asyncio
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from config.database.mysql import get_engine, get_session_factory
|
||||
from repositories.sys_config import SysConfigRepo
|
||||
|
||||
|
||||
async def _db_ok() -> bool:
|
||||
"""MySQL 可达且 sys_config 表已建(schema.sql 已导入)才算就绪。"""
|
||||
try:
|
||||
async with get_engine().connect() as conn:
|
||||
await conn.execute(text("SELECT 1"))
|
||||
async with get_session_factory()() as session:
|
||||
await session.execute(text("SELECT 1 FROM sys_config LIMIT 1"))
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
async def main():
|
||||
try:
|
||||
if not await _db_ok():
|
||||
print("SKIP: MySQL 不可达或 sys_config 未建表,CRUD 检查延后(先启动 MySQL 并导入 sql/schema.sql)")
|
||||
return
|
||||
async with get_session_factory()() as session:
|
||||
repo = SysConfigRepo(session)
|
||||
row = await repo.set_value("test.hnw", "3000000", "自检")
|
||||
assert row.config_key == "test.hnw"
|
||||
assert await repo.get_value("test.hnw") == "3000000"
|
||||
assert await repo.count() > 0
|
||||
rows = await repo.list(where=(repo.model.config_key == "test.hnw",))
|
||||
assert len(rows) == 1
|
||||
assert await repo.delete(rows[0].id)
|
||||
assert await repo.get_value("test.hnw") is None # 幂等清理
|
||||
print("REPO TESTS PASSED")
|
||||
finally:
|
||||
import config.database as _db
|
||||
await _db.dispose() # 干净收尾,避免连接在事件循环关闭后被 GC
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
Reference in New Issue
Block a user