Files
llm-in-text/backend/main.py
T

485 lines
15 KiB
Python
Raw Normal View History

import base64
import json
import logging
import os
import uuid
from typing import Optional
from fastapi import FastAPI, HTTPException, Request, Security
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, StreamingResponse
from fastapi.security import APIKeyHeader
from pydantic import BaseModel
from geoip import get_ip_location_text
2026-06-06 15:44:00 +08:00
from job_handlers import (
_sanitize_converted_markdown,
sanitize_inline_completion_content,
ALLOWED_CONVERT_EXTENSIONS,
asr_handler,
completion_handler,
compress_handler,
convert_handler,
ocr_handler,
pro_completion_handler,
tts_handler,
)
from job_system import (
InMemoryJobManager,
JobSystemError,
JOB_TYPES,
QueueFullError,
RedisJobManager,
get_job_manager,
persist_temp_input,
)
from models import UserPreferences
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s - %(message)s",
)
logger = logging.getLogger("api")
app = FastAPI()
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*", "X-API-Key", "X-Client-IP", "X-Request-Id"],
)
API_KEY = os.getenv("API_KEY", "your-secret-key-here")
2026-06-06 15:44:00 +08:00
DOC_COMPRESS_CONTEXT_LIMIT = int(os.getenv("DOC_COMPRESS_CONTEXT_LIMIT", "128000"))
api_key_header = APIKeyHeader(name="X-API-Key")
2026-06-06 15:44:00 +08:00
_handlers_registered = False
class CompletionRequest(BaseModel):
prefix: str
suffix: str
languageId: str = "markdown"
model_thinking: str = "low"
privacy_mode: bool = False
user_preferences: Optional[UserPreferences] = None
model: Optional[str] = None
temperature: float = 0.7
2026-06-06 15:44:00 +08:00
class ProCompletionRequest(BaseModel):
prefix: str
suffix: str
languageId: str = "markdown"
instruction: str = ""
pro_thinking: str = "medium"
privacy_mode: bool = False
user_preferences: Optional[UserPreferences] = None
class CancelCompletionRequest(BaseModel):
request_id: str
reason: str = "abort"
class OCRRequest(BaseModel):
image: str
filename: str = "image.jpg"
language: str = "auto"
class ConvertRequest(BaseModel):
file: str
filename: str = "document.pdf"
2026-06-06 15:44:00 +08:00
class CompressRequest(BaseModel):
content: str
docType: str = "txt"
2026-06-06 15:44:00 +08:00
class TTSJobRequest(BaseModel):
text: str
instruct: str = ""
speaker: str = "Vivian"
format: str = "wav"
2026-06-06 15:44:00 +08:00
class ASRJobRequest(BaseModel):
audio_base64: str
language: Optional[str] = "zh-CN"
def _preview(text: str, limit: int = 80) -> str:
value = (text or "").replace("\n", "\\n")
if len(value) <= limit:
return value
return value[:limit] + "..."
def get_client_ip(request: Request) -> str:
if request.client:
return request.headers.get("X-Client-IP") or request.client.host
return request.headers.get("X-Client-IP") or "unknown"
def _clamp_temperature(value: float, default: float = 0.7) -> float:
try:
numeric = float(value)
except (TypeError, ValueError):
return default
return max(0.0, min(numeric, 1.2))
2026-06-06 15:44:00 +08:00
async def get_api_key(api_key: str = Security(api_key_header)): # pragma: no cover
if api_key != API_KEY:
raise HTTPException(status_code=403, detail="Could not validate credentials")
return api_key
2026-06-06 15:44:00 +08:00
def _serialize_preferences(preferences: UserPreferences | None) -> dict | None:
if preferences is None:
return None
if hasattr(preferences, "dict"):
return preferences.dict()
return dict(preferences)
2026-06-06 15:44:00 +08:00
def _request_id(request: Request) -> str:
return request.headers.get("X-Request-Id") or str(uuid.uuid4())
2026-06-06 15:44:00 +08:00
def _register_handlers() -> None:
global _handlers_registered
if _handlers_registered:
return
manager = get_job_manager()
manager.register_handler("completion", completion_handler)
manager.register_handler("pro_completion", pro_completion_handler)
manager.register_handler("compress", compress_handler)
manager.register_handler("ocr", ocr_handler)
manager.register_handler("convert", convert_handler)
manager.register_handler("tts", tts_handler)
manager.register_handler("asr", asr_handler)
_handlers_registered = True
2026-06-06 15:44:00 +08:00
def _sse(event: str, data: dict) -> str:
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
2026-06-06 15:44:00 +08:00
async def _stream_job(job_id: str):
_register_handlers()
manager = get_job_manager()
async def event_stream():
try:
2026-06-06 15:44:00 +08:00
async for event in manager.stream_events(job_id):
event_name = event.get("event", "message")
payload = {k: v for k, v in event.items() if k != "event"}
yield _sse(event_name, payload)
except Exception as exc:
logger.exception("job stream failed job_id=%s", job_id)
yield _sse("error", {"job_id": job_id, "error": str(exc)})
return StreamingResponse(
event_stream(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)
2026-06-06 15:44:00 +08:00
async def _queue_job(job_type: str, payload: dict, request_id: str) -> str:
_register_handlers()
manager = get_job_manager()
return await manager.submit(job_type, payload, request_id=request_id)
async def _cancel_job(request_id: str, reason: str) -> dict:
_register_handlers()
manager = get_job_manager()
return await manager.cancel(request_id, reason)
async def _job_status(job_id: str) -> dict | None:
_register_handlers()
manager = get_job_manager()
return await manager.get_status(job_id)
async def _queue_load_snapshot() -> dict:
manager = get_job_manager()
if isinstance(manager, InMemoryJobManager):
return {job_type: manager._metrics(job_type) for job_type in manager.queues}
if isinstance(manager, RedisJobManager):
return {job_type: await manager._metrics(job_type) for job_type in JOB_TYPES}
return {}
@app.post("/v1/completions")
async def create_completion(
request: Request,
req: CompletionRequest,
api_key: str = Security(get_api_key),
):
del api_key
request_id = _request_id(request)
location = ""
if not req.privacy_mode: # pragma: no cover
location = get_ip_location_text(get_client_ip(request))
payload = {
"request_id": request_id,
"location": location,
"request": {
"prefix": req.prefix,
"suffix": req.suffix,
"languageId": req.languageId,
"model_thinking": req.model_thinking,
"privacy_mode": req.privacy_mode,
"user_preferences": _serialize_preferences(req.user_preferences),
"model": req.model,
"temperature": _clamp_temperature(req.temperature, 0.7),
},
}
try:
job_id = await _queue_job("completion", payload, request_id)
except QueueFullError as exc:
return JSONResponse({"error": str(exc), "request_id": request_id}, status_code=429)
except JobSystemError as exc:
return JSONResponse({"error": str(exc), "request_id": request_id}, status_code=503)
return await _stream_job(job_id)
@app.post("/v1/completions/cancel")
async def cancel_completion(req: CancelCompletionRequest, api_key: str = Security(get_api_key)):
2026-06-06 15:44:00 +08:00
del api_key
return await _cancel_job(req.request_id or "", req.reason)
@app.post("/v1/pro/completions")
async def create_pro_completion(
request: Request,
req: ProCompletionRequest,
api_key: str = Security(get_api_key),
):
del api_key
request_id = _request_id(request)
location = ""
if not req.privacy_mode: # pragma: no cover
location = get_ip_location_text(get_client_ip(request))
payload = {
"request_id": request_id,
"location": location,
"request": {
"prefix": req.prefix,
"suffix": req.suffix,
"languageId": req.languageId,
"instruction": req.instruction,
"pro_thinking": req.pro_thinking,
"privacy_mode": req.privacy_mode,
"user_preferences": _serialize_preferences(req.user_preferences),
},
}
try:
job_id = await _queue_job("pro_completion", payload, request_id)
except QueueFullError as exc:
return JSONResponse({"error": str(exc), "request_id": request_id}, status_code=429)
except JobSystemError as exc:
return JSONResponse({"error": str(exc), "request_id": request_id}, status_code=503)
return await _stream_job(job_id)
@app.post("/v1/pro/completions/cancel")
async def cancel_pro_completion(req: CancelCompletionRequest, api_key: str = Security(get_api_key)):
del api_key
return await _cancel_job(req.request_id or "", req.reason)
@app.get("/v1/pro/completions/status/{request_id}")
async def get_pro_completion_status(request_id: str, api_key: str = Security(get_api_key)):
del api_key
state = await _job_status(request_id)
if state is None:
raise HTTPException(status_code=404, detail="PRO request not found")
return state
@app.post("/v1/ocr")
2026-06-06 15:44:00 +08:00
async def ocr_image(req: OCRRequest, api_key: str = Security(get_api_key)):
del api_key
request_id = str(uuid.uuid4())
try:
2026-06-06 15:44:00 +08:00
image_bytes = base64.b64decode(req.image)
except Exception as exc:
return JSONResponse({"error": str(exc)}, status_code=500)
input_path = persist_temp_input(image_bytes, os.path.splitext(req.filename)[1] or ".img")
try:
job_id = await _queue_job("ocr", {
"request_id": request_id,
"input_path": input_path,
"filename": req.filename,
"language": req.language,
}, request_id)
except Exception:
if os.path.exists(input_path):
os.unlink(input_path)
raise
return await _stream_job(job_id)
@app.post("/v1/convert")
2026-06-06 15:44:00 +08:00
async def convert_to_markdown(req: ConvertRequest, api_key: str = Security(get_api_key)):
del api_key
request_id = str(uuid.uuid4())
ext = os.path.splitext(req.filename)[1].lower()
if ext not in ALLOWED_CONVERT_EXTENSIONS:
return JSONResponse({"error": "仅支持 txt、docx、pptx、pdf 格式"}, status_code=500)
try:
2026-06-06 15:44:00 +08:00
file_bytes = base64.b64decode(req.file)
except Exception as exc:
return JSONResponse({"error": str(exc)}, status_code=500)
input_path = persist_temp_input(file_bytes, ext or ".bin")
try:
job_id = await _queue_job("convert", {
"request_id": request_id,
"input_path": input_path,
"filename": req.filename,
}, request_id)
except Exception:
if os.path.exists(input_path):
os.unlink(input_path)
raise
return await _stream_job(job_id)
@app.post("/v1/compress/submit")
async def submit_compress(req: CompressRequest, api_key: str = Security(get_api_key)):
del api_key
content = req.content or ""
if not content.strip():
raise HTTPException(status_code=400, detail="文档内容为空,无法压缩")
if len(content) > DOC_COMPRESS_CONTEXT_LIMIT:
raise HTTPException(
status_code=400,
detail=f"文档内容过长({len(content)} 字符),超过限制 {DOC_COMPRESS_CONTEXT_LIMIT},无法压缩",
)
2026-06-06 15:44:00 +08:00
task_id = str(uuid.uuid4())
await _queue_job("compress", {"request_id": task_id, "content": content, "docType": req.docType or "txt"}, task_id)
return {"task_id": task_id, "status": "queued"}
@app.get("/v1/compress/status")
async def get_compress_status(task_id: str, api_key: str = Security(get_api_key)):
del api_key
if not task_id:
raise HTTPException(status_code=400, detail="缺少 task_id 参数")
state = await _job_status(task_id)
if state is None:
raise HTTPException(status_code=404, detail="任务不存在或已过期")
if state["status"] == "completed":
result = state.get("result") or {}
return {"task_id": task_id, "status": "completed", "content": result.get("content", "")}
if state["status"] == "failed":
return {"task_id": task_id, "status": "error", "message": state.get("error") or ""}
if state["status"] == "cancelled":
return {"task_id": task_id, "status": "cancelled"}
return {
"task_id": task_id,
"status": "processing" if state["status"] == "running" else "queued",
"busy_level": state.get("busy_level"),
"queued_count": state.get("queued_count"),
"running_count": state.get("running_count"),
}
@app.post("/v1/tts-asr/tts")
async def queue_tts(req: TTSJobRequest, request: Request, api_key: str = Security(get_api_key)):
del api_key
request_id = _request_id(request)
job_id = await _queue_job("tts", {
"request_id": request_id,
"text": req.text,
"instruct": req.instruct,
"speaker": req.speaker,
"format": req.format,
}, request_id)
return await _stream_job(job_id)
@app.post("/v1/tts-asr/asr")
async def queue_asr(req: ASRJobRequest, request: Request, api_key: str = Security(get_api_key)):
del api_key
request_id = _request_id(request)
try:
audio_bytes = base64.b64decode(req.audio_base64)
except Exception as exc:
return JSONResponse({"error": str(exc)}, status_code=500)
input_path = persist_temp_input(audio_bytes, ".wav")
try:
job_id = await _queue_job("asr", {
"request_id": request_id,
"input_path": input_path,
"language": req.language or "zh-CN",
}, request_id)
except Exception:
if os.path.exists(input_path):
os.unlink(input_path)
raise
return await _stream_job(job_id)
2026-06-06 15:44:00 +08:00
@app.post("/v1/jobs/{job_id}/cancel")
async def cancel_job(job_id: str, req: CancelCompletionRequest, api_key: str = Security(get_api_key)):
del api_key
return await _cancel_job(req.request_id or job_id, req.reason)
2026-06-06 15:44:00 +08:00
@app.get("/v1/jobs/{job_id}/status")
async def get_job_status(job_id: str, api_key: str = Security(get_api_key)):
del api_key
state = await _job_status(job_id)
if state is None:
raise HTTPException(status_code=404, detail="job not found")
return state
@app.get("/v1/jobs/load")
async def get_job_load(api_key: str = Security(get_api_key)):
del api_key
return {"queues": await _queue_load_snapshot()}
def _register_tts_asr_routes():
try:
from tts_asr import register_tts_asr_routes
except ModuleNotFoundError as exc:
logger.warning("Skipping TTS/ASR route registration because a dependency is missing: %s", exc)
return
except Exception as exc:
logger.warning("Skipping TTS/ASR route registration because import failed: %s", exc)
return
try:
2026-06-06 15:44:00 +08:00
register_tts_asr_routes(app, include_generation_routes=False)
except Exception as exc:
logger.warning("Failed to register TTS/ASR routes: %s", exc)
2026-06-06 15:44:00 +08:00
_register_tts_asr_routes()
2026-06-06 15:44:00 +08:00
@app.on_event("shutdown")
async def _shutdown_job_manager(): # pragma: no cover
manager = get_job_manager()
close = getattr(manager, "close", None)
if close is not None:
await close()
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8001)