fix: 核心 bug 修复 - API Key 加密/解密 + AI 调用 + 数据安全
- schema: 用 @field_serializer 掩码 api_key,不再直接修改 SQLAlchemy 对象 - API: 移除所有 model.api_key='***' 赋值,防止数据库被覆盖 - security: SHA256(SECRET_KEY) 派生合法 Fernet 32 字节密钥 - ai_adapter: CustomAdapter 改用 OpenAI-compatible messages 格式 - ai_adapter: trust_env=False 防止代理干扰,增加详细异常信息 - templates: bleach HTML 净化防止 XSS - .env.example: CORS_ORIGINS 修复为 JSON 数组格式 - 端到端验证:单点测试 + 异步生成任务均通过
This commit is contained in:
parent
aeff8753ec
commit
40b861cfcc
18
.env.example
18
.env.example
|
|
@ -2,9 +2,25 @@
|
||||||
POSTGRES_USER=docforge
|
POSTGRES_USER=docforge
|
||||||
POSTGRES_PASSWORD=change-this-password
|
POSTGRES_PASSWORD=change-this-password
|
||||||
POSTGRES_DB=docforge
|
POSTGRES_DB=docforge
|
||||||
|
POSTGRES_HOST=localhost
|
||||||
|
POSTGRES_PORT=5432
|
||||||
|
|
||||||
# Security (generate with: python -c "import secrets; print(secrets.token_urlsafe(32))")
|
# Redis
|
||||||
|
REDIS_URL=redis://localhost:6379/0
|
||||||
|
|
||||||
|
# Backend
|
||||||
SECRET_KEY=change-me-to-a-random-secret-key
|
SECRET_KEY=change-me-to-a-random-secret-key
|
||||||
|
DEFAULT_MODEL_ID=
|
||||||
|
|
||||||
|
# File Storage
|
||||||
|
STORAGE_ROOT=./storage
|
||||||
|
|
||||||
|
# Celery
|
||||||
|
CELERY_BROKER_URL=redis://localhost:6379/0
|
||||||
|
CELERY_RESULT_BACKEND=redis://localhost:6379/1
|
||||||
|
|
||||||
|
# CORS (JSON array format)
|
||||||
|
CORS_ORIGINS=["http://localhost:3000","http://localhost:5173"]
|
||||||
|
|
||||||
# AI (optional, can be configured in the web UI)
|
# AI (optional, can be configured in the web UI)
|
||||||
OPENAI_API_KEY=
|
OPENAI_API_KEY=
|
||||||
|
|
|
||||||
|
|
@ -24,7 +24,6 @@ async def create_model(data: AIModelCreate, db: AsyncSession = Depends(get_db)):
|
||||||
db.add(model)
|
db.add(model)
|
||||||
await db.flush()
|
await db.flush()
|
||||||
await db.refresh(model)
|
await db.refresh(model)
|
||||||
model.api_key = "***"
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -41,8 +40,6 @@ async def list_models(
|
||||||
query = query.offset(skip).limit(limit).order_by(AIModel.created_at.desc())
|
query = query.offset(skip).limit(limit).order_by(AIModel.created_at.desc())
|
||||||
result = await db.execute(query)
|
result = await db.execute(query)
|
||||||
models = result.scalars().all()
|
models = result.scalars().all()
|
||||||
for m in models:
|
|
||||||
m.api_key = "***"
|
|
||||||
return models
|
return models
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -52,7 +49,6 @@ async def get_model(model_id: str, db: AsyncSession = Depends(get_db)):
|
||||||
model = result.scalar_one_or_none()
|
model = result.scalar_one_or_none()
|
||||||
if not model:
|
if not model:
|
||||||
raise HTTPException(status_code=404, detail="模型不存在")
|
raise HTTPException(status_code=404, detail="模型不存在")
|
||||||
model.api_key = "***"
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -72,7 +68,6 @@ async def update_model(model_id: str, data: AIModelUpdate, db: AsyncSession = De
|
||||||
|
|
||||||
await db.flush()
|
await db.flush()
|
||||||
await db.refresh(model)
|
await db.refresh(model)
|
||||||
model.api_key = "***"
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -95,5 +90,4 @@ async def toggle_model(model_id: str, data: AIModelToggle, db: AsyncSession = De
|
||||||
model.is_enabled = data.is_enabled
|
model.is_enabled = data.is_enabled
|
||||||
await db.flush()
|
await db.flush()
|
||||||
await db.refresh(model)
|
await db.refresh(model)
|
||||||
model.api_key = "***"
|
|
||||||
return model
|
return model
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
import uuid
|
import uuid
|
||||||
|
import bleach
|
||||||
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form, Query
|
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form, Query
|
||||||
from fastapi.responses import Response
|
from fastapi.responses import Response
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
@ -29,8 +30,6 @@ ALLOWED_ATTRS = {
|
||||||
"table": ["style"],
|
"table": ["style"],
|
||||||
}
|
}
|
||||||
|
|
||||||
import bleach
|
|
||||||
|
|
||||||
router = APIRouter(prefix="/templates", tags=["模板管理"])
|
router = APIRouter(prefix="/templates", tags=["模板管理"])
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field, field_serializer
|
||||||
|
|
||||||
|
|
||||||
class AIModelCreate(BaseModel):
|
class AIModelCreate(BaseModel):
|
||||||
|
|
@ -32,7 +32,7 @@ class AIModelResponse(BaseModel):
|
||||||
name: str
|
name: str
|
||||||
provider: str
|
provider: str
|
||||||
endpoint: str
|
endpoint: str
|
||||||
api_key: str = Field(default="***")
|
api_key: str
|
||||||
extra_params: dict
|
extra_params: dict
|
||||||
is_enabled: bool
|
is_enabled: bool
|
||||||
remark: str | None = None
|
remark: str | None = None
|
||||||
|
|
@ -40,3 +40,7 @@ class AIModelResponse(BaseModel):
|
||||||
updated_at: datetime
|
updated_at: datetime
|
||||||
|
|
||||||
model_config = {"from_attributes": True}
|
model_config = {"from_attributes": True}
|
||||||
|
|
||||||
|
@field_serializer("api_key")
|
||||||
|
def mask_api_key(self, value: str) -> str:
|
||||||
|
return "***"
|
||||||
|
|
|
||||||
|
|
@ -105,12 +105,16 @@ class CustomAdapter(AIAdapter):
|
||||||
if ref_content:
|
if ref_content:
|
||||||
user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}"
|
user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}"
|
||||||
|
|
||||||
|
messages = [{"role": "system", "content": "你是一个专业的文档内容生成助手。"}]
|
||||||
|
messages.append({"role": "user", "content": user_content})
|
||||||
|
|
||||||
body = {
|
body = {
|
||||||
"prompt": user_content,
|
"model": extra.get("model", "gpt-3.5-turbo"),
|
||||||
|
"messages": messages,
|
||||||
"max_tokens": extra.get("max_tokens", 2000),
|
"max_tokens": extra.get("max_tokens", 2000),
|
||||||
"temperature": extra.get("temperature", 0.7),
|
"temperature": extra.get("temperature", 0.7),
|
||||||
}
|
}
|
||||||
body.update({k: v for k, v in extra.items() if k not in ("max_tokens", "temperature")})
|
body.update({k: v for k, v in extra.items() if k not in ("model", "messages", "max_tokens", "temperature")})
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"url": model_config["endpoint"],
|
"url": model_config["endpoint"],
|
||||||
|
|
@ -153,11 +157,25 @@ async def call_ai_model(model_config: dict, prompt: str, ref_content: str | None
|
||||||
|
|
||||||
timeout = model_config.get("extra_params", {}).get("timeout", 120)
|
timeout = model_config.get("extra_params", {}).get("timeout", 120)
|
||||||
|
|
||||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
try:
|
||||||
response = await client.post(
|
async with httpx.AsyncClient(
|
||||||
request["url"],
|
timeout=timeout,
|
||||||
headers=request["headers"],
|
proxy=None,
|
||||||
json=request["json"],
|
trust_env=False,
|
||||||
)
|
) as client:
|
||||||
response.raise_for_status()
|
response = await client.post(
|
||||||
return adapter.parse_response(response.json())
|
request["url"],
|
||||||
|
headers=request["headers"],
|
||||||
|
json=request["json"],
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
return adapter.parse_response(response.json())
|
||||||
|
except httpx.HTTPStatusError as e:
|
||||||
|
detail = e.response.text[:500] if e.response else str(e)
|
||||||
|
raise RuntimeError(f"AI 服务返回错误 ({e.response.status_code}): {detail}")
|
||||||
|
except httpx.TimeoutException:
|
||||||
|
raise RuntimeError("AI 调用超时")
|
||||||
|
except httpx.ConnectError as e:
|
||||||
|
raise RuntimeError(f"无法连接 AI 服务: {e}")
|
||||||
|
except Exception as e:
|
||||||
|
raise RuntimeError(f"AI 调用异常: {str(e)}")
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue