NestJS (6 services): implement @RequirePermission decorator with SetMetadata+Reflector, register APP_GUARD globally, fix as assertions to type guards, add explicit return types, fix import type for express, fix /metrics implicit any, replace native Error with ApplicationError, remove typeorm remnants, register LifecycleService. teacher-bff: add logger, ApplicationError, GlobalErrorFilter, forward real userId to downstream, log downstream failures, migrate health controller to shared/health. Go (2 services): interface to any, doc comments, CORS dev whitelist, JWT secret fail-fast, push-gateway internal API auth, metrics and readyz endpoints, remove dead code. Python (2 services): lifespan return type, dev_mode to bool, data-ana APIRouter, ai POST body model, ClickHouse async wrapping.
248 lines
7.7 KiB
Python
248 lines
7.7 KiB
Python
"""AI 网关服务入口."""
|
||
|
||
from collections.abc import AsyncGenerator
|
||
from contextlib import asynccontextmanager
|
||
from typing import Any
|
||
|
||
import structlog
|
||
from fastapi import APIRouter, FastAPI
|
||
from fastapi.responses import StreamingResponse
|
||
from opentelemetry import trace
|
||
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
|
||
from opentelemetry.instrumentation.fastapi import FastAPIInstrumentor
|
||
from opentelemetry.sdk.trace import TracerProvider
|
||
from opentelemetry.sdk.trace.export import BatchSpanProcessor
|
||
from prometheus_client import make_asgi_app
|
||
from pydantic import BaseModel
|
||
|
||
from .config import settings
|
||
from .llm_client import chat_completion, chat_completion_stream
|
||
|
||
logger = structlog.get_logger()
|
||
tracer = trace.get_tracer(__name__)
|
||
|
||
|
||
def init_tracer() -> None:
|
||
"""初始化 OpenTelemetry.
|
||
|
||
endpoint 从 settings.otel_endpoint 读取;dev_mode=true 时跳过 exporter
|
||
初始化,避免本地无 collector 时报错。
|
||
"""
|
||
if settings.is_dev:
|
||
logger.info("dev_mode_tracer_skipped", dev_mode=settings.dev_mode)
|
||
return
|
||
|
||
provider = TracerProvider()
|
||
endpoint = f"{settings.otel_endpoint.rstrip('/')}/v1/traces"
|
||
exporter = OTLPSpanExporter(endpoint=endpoint)
|
||
provider.add_span_processor(BatchSpanProcessor(exporter))
|
||
trace.set_tracer_provider(provider)
|
||
logger.info("tracer_initialized", otel_endpoint=endpoint)
|
||
|
||
|
||
@asynccontextmanager
|
||
async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
|
||
"""应用生命周期."""
|
||
init_tracer()
|
||
logger.info(
|
||
"ai_service_starting",
|
||
llm_available=settings.llm_available,
|
||
dev_mode=settings.is_dev,
|
||
openai_base_url=settings.openai_base_url,
|
||
)
|
||
if not settings.llm_available:
|
||
logger.warning("ai_service_llm_degraded_no_api_key")
|
||
yield
|
||
logger.info("ai_service_stopping")
|
||
|
||
|
||
app = FastAPI(
|
||
title="AI Gateway Service",
|
||
version="0.1.0",
|
||
lifespan=lifespan,
|
||
)
|
||
|
||
# OpenTelemetry FastAPI 自动埋点(HTTP 请求/响应 span)
|
||
FastAPIInstrumentor.instrument_app(app)
|
||
|
||
app.mount("/metrics", make_asgi_app())
|
||
|
||
# 业务路由加 /ai 前缀,Gateway 代理 /api/v1/ai/* → /ai/*
|
||
router = APIRouter(prefix="/ai")
|
||
|
||
|
||
class ChatRequest(BaseModel):
|
||
"""聊天请求."""
|
||
|
||
messages: list[dict[str, Any]]
|
||
model: str = "gpt-4o-mini"
|
||
temperature: float = 0.7
|
||
stream: bool = False
|
||
|
||
|
||
class ChatResponse(BaseModel):
|
||
"""聊天响应."""
|
||
|
||
content: str
|
||
model: str
|
||
usage: dict[str, Any]
|
||
degraded: bool = False
|
||
|
||
|
||
class QuestionRequest(BaseModel):
|
||
"""题目生成请求."""
|
||
|
||
prompt: str
|
||
|
||
|
||
def _extract_content(result: dict[str, Any] | None) -> tuple[str, str, dict[str, Any]]:
|
||
"""从 OpenAI 响应中抽取 (content, model, usage)。"""
|
||
if result is None:
|
||
return "", "", {}
|
||
choices = result.get("choices", [])
|
||
content = ""
|
||
if choices:
|
||
content = choices[0].get("message", {}).get("content", "") or ""
|
||
model = result.get("model", "") or ""
|
||
usage = result.get("usage", {}) or {}
|
||
return content, model, usage
|
||
|
||
|
||
@app.get("/healthz")
|
||
async def healthz() -> dict[str, Any]:
|
||
"""健康检查(liveness)."""
|
||
return {"status": "ok", "service": "ai"}
|
||
|
||
|
||
@app.get("/readyz")
|
||
async def readyz() -> dict[str, Any]:
|
||
"""就绪检查(readiness).
|
||
|
||
LLM 未配置时仍返回 200,但标记 degraded=true,调用方可据此判断是否路由流量。
|
||
"""
|
||
llm_configured = settings.llm_available
|
||
return {
|
||
"status": "ok",
|
||
"service": "ai",
|
||
"llm_configured": llm_configured,
|
||
"degraded": not llm_configured,
|
||
"openai_base_url": settings.openai_base_url,
|
||
}
|
||
|
||
|
||
@router.post("/chat", response_model=ChatResponse)
|
||
async def chat(req: ChatRequest) -> ChatResponse:
|
||
"""LLM 聊天接口(无 API key 时降级返回骨架响应)."""
|
||
with tracer.start_as_current_span("ai_chat"):
|
||
result = await chat_completion(
|
||
messages=req.messages,
|
||
model=req.model,
|
||
temperature=req.temperature,
|
||
api_key=settings.openai_api_key,
|
||
base_url=settings.openai_base_url,
|
||
)
|
||
if result is None:
|
||
logger.warning("chat_degraded", model=req.model)
|
||
return ChatResponse(
|
||
content="[degraded] LLM unavailable - returning skeleton response",
|
||
model=req.model,
|
||
usage={"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0},
|
||
degraded=True,
|
||
)
|
||
content, model, usage = _extract_content(result)
|
||
return ChatResponse(
|
||
content=content,
|
||
model=model or req.model,
|
||
usage=usage,
|
||
degraded=False,
|
||
)
|
||
|
||
|
||
@router.post("/chat/stream")
|
||
async def chat_stream(req: ChatRequest) -> StreamingResponse:
|
||
"""流式聊天(SSE,无 API key 时降级返回骨架 SSE)."""
|
||
|
||
async def generate() -> AsyncGenerator[str, None]:
|
||
with tracer.start_as_current_span("ai_chat_stream"):
|
||
async for chunk in chat_completion_stream(
|
||
messages=req.messages,
|
||
model=req.model,
|
||
temperature=req.temperature,
|
||
api_key=settings.openai_api_key,
|
||
base_url=settings.openai_base_url,
|
||
):
|
||
yield chunk
|
||
|
||
return StreamingResponse(generate(), media_type="text/event-stream")
|
||
|
||
|
||
@router.post("/generate/question")
|
||
async def generate_question(req: QuestionRequest) -> dict[str, Any]:
|
||
"""生成题目(无 API key 时降级返回骨架)."""
|
||
with tracer.start_as_current_span("generate_question"):
|
||
messages = [
|
||
{
|
||
"role": "system",
|
||
"content": "You are an educational question generator. "
|
||
"Generate a clear, concise question based on the user's prompt.",
|
||
},
|
||
{"role": "user", "content": req.prompt},
|
||
]
|
||
result = await chat_completion(
|
||
messages=messages,
|
||
model="gpt-4o-mini",
|
||
temperature=0.7,
|
||
api_key=settings.openai_api_key,
|
||
base_url=settings.openai_base_url,
|
||
)
|
||
if result is None:
|
||
logger.warning("generate_question_degraded", prompt=req.prompt[:100])
|
||
return {
|
||
"success": True,
|
||
"data": {"question": "[degraded] question generation skeleton"},
|
||
"degraded": True,
|
||
}
|
||
content, _, _ = _extract_content(result)
|
||
return {
|
||
"success": True,
|
||
"data": {"question": content},
|
||
"degraded": False,
|
||
}
|
||
|
||
|
||
@router.post("/optimize/expression")
|
||
async def optimize_expression(text: str) -> dict[str, Any]:
|
||
"""优化表达(无 API key 时降级返回骨架)."""
|
||
with tracer.start_as_current_span("optimize_expression"):
|
||
messages = [
|
||
{
|
||
"role": "system",
|
||
"content": "You are a writing assistant. "
|
||
"Optimize the user's text for clarity, conciseness, and tone.",
|
||
},
|
||
{"role": "user", "content": text},
|
||
]
|
||
result = await chat_completion(
|
||
messages=messages,
|
||
model="gpt-4o-mini",
|
||
temperature=0.5,
|
||
api_key=settings.openai_api_key,
|
||
base_url=settings.openai_base_url,
|
||
)
|
||
if result is None:
|
||
logger.warning("optimize_expression_degraded", text=text[:100])
|
||
return {
|
||
"success": True,
|
||
"data": {"optimized": "[degraded] expression optimization skeleton"},
|
||
"degraded": True,
|
||
}
|
||
content, _, _ = _extract_content(result)
|
||
return {
|
||
"success": True,
|
||
"data": {"optimized": content},
|
||
"degraded": False,
|
||
}
|
||
|
||
|
||
app.include_router(router)
|