第一次提交

This commit is contained in:
2026-08-01 16:50:55 +08:00
commit c8f9e39039
68 changed files with 6016 additions and 0 deletions
View File
+24
View File
@@ -0,0 +1,24 @@
"""用户 AI api_key 的对称加密(Fernet)。密钥来自 settings.fernet_key。"""
from cryptography.fernet import Fernet, InvalidToken
from app.config import settings
def _fernet() -> Fernet:
if not settings.fernet_key:
raise RuntimeError(
"FERNET_KEY 未配置,无法加解密用户 api_key。"
"生成:python -c \"from cryptography.fernet import Fernet; print(Fernet.generate_key().decode())\""
)
return Fernet(settings.fernet_key.encode("utf-8"))
def encrypt(plaintext: str) -> str:
return _fernet().encrypt(plaintext.encode("utf-8")).decode("utf-8")
def decrypt(token: str) -> str | None:
try:
return _fernet().decrypt(token.encode("utf-8")).decode("utf-8")
except (InvalidToken, ValueError):
return None
+166
View File
@@ -0,0 +1,166 @@
"""上传图片 → OCR → 结构化提取 → 生成草稿题目。作为 BackgroundTask 运行。
注意:BackgroundTask 在请求返回后执行,需要自己开新的 DB session。
"""
import json
import logging
from pydantic import ValidationError
from sqlalchemy import select
from app.db import SessionLocal
from app.models.image import Image
from app.models.job import JobStatus, ProcessingJob
from app.models.question import Question, QuestionType
from app.models.user import UserSettings
from app.schemas.question import QuestionCreate
from app.services.llm.client import LLMClient, LLMConfigError
from app.services.llm.prompts import EXTRACT_SYSTEM, EXTRACT_USER_TEMPLATE
from app.services.ocr.pipeline import run_ocr
from app.services.storage import download_image_async
logger = logging.getLogger(__name__)
def _parse_extract_json(raw: str) -> dict:
"""从模型输出里解析 JSON(容忍代码块包裹)。"""
text = raw.strip()
if text.startswith("```"):
text = text.split("```", 2)[1]
if text.startswith("json"):
text = text[4:]
return json.loads(text.strip())
async def run_extract_job(job_id: str, user_id: str, object_key: str, mime: str):
async with SessionLocal() as db:
job = await db.get(ProcessingJob, job_id)
if job is None:
logger.error("job %s 不存在", job_id)
return
try:
# 加载用户 AI 配置
settings_row = await db.scalar(
select(UserSettings).where(UserSettings.user_id == user_id)
)
llm = LLMClient(settings_row)
# 1. OCR
job.status = JobStatus.ocr_running
await db.commit()
image_bytes = await download_image_async(object_key)
ocr = await run_ocr(image_bytes, mime, llm)
# 2. 结构化提取
job.status = JobStatus.ai_running
job.engine_used = ocr.engine
await db.commit()
extract = await llm.chat_text(
EXTRACT_SYSTEM,
EXTRACT_USER_TEMPLATE.format(content=ocr.markdown),
json_mode=True,
)
questions = _build_questions(user_id, extract.content, ocr, job.image_id)
# 缓存原始 OCR 文本,便于日后追溯或重新拆分
if job.image_id:
image = await db.get(Image, job.image_id)
if image is not None:
image.ocr_markdown = ocr.markdown
db.add_all(questions)
await db.flush()
job.result_question_ids = [q.id for q in questions]
job.status = JobStatus.done
await db.commit()
logger.info(
"job %s 完成,engine=%s,拆出 %d 条笔记",
job_id,
ocr.engine,
len(questions),
)
except LLMConfigError as e:
await _fail(db, job, f"AI 未配置:{e}")
except Exception as e: # noqa: BLE001
logger.exception("job %s 失败", job_id)
await _fail(db, job, str(e))
def _build_questions(
user_id: str, extract_raw: str, ocr, image_id: str | None
) -> list[Question]:
"""把提取结果拆成多条笔记。
一张照片常含多道题,每道题存成独立笔记。整体解析失败时不丢内容,
降级为一条 unclassified 笔记(原文进题干),用户可自行整理。
"""
try:
data = _parse_extract_json(extract_raw)
raw_items = data.get("questions")
if not isinstance(raw_items, list) or not raw_items:
raise ValueError("提取结果里没有 questions 数组")
except (json.JSONDecodeError, KeyError, ValueError) as e:
logger.warning("提取结果解析失败,降级为一条未分类笔记:%s", e)
return [_fallback_question(user_id, ocr, image_id)]
questions: list[Question] = []
for idx, item in enumerate(raw_items):
try:
questions.append(_build_one(user_id, item, ocr, image_id))
except (KeyError, ValueError, ValidationError, TypeError) as e:
# 单条坏掉不影响其他条:这条退化为未分类,原文留在题干里
logger.warning("%d 条解析失败,存为未分类:%s", idx + 1, e)
stem = ""
if isinstance(item, dict):
stem = str(item.get("stem_markdown") or "")
questions.append(
_fallback_question(user_id, ocr, image_id, stem_override=stem or None)
)
return questions or [_fallback_question(user_id, ocr, image_id)]
def _build_one(user_id: str, item: dict, ocr, image_id: str | None) -> Question:
parsed = QuestionCreate(
# 模型没给或给了不认识的题型,就当未分类,别硬猜
type=_coerce_type(item.get("type")),
stem_markdown=item.get("stem_markdown") or ocr.markdown,
options=item.get("options"),
correct_answer=item.get("correct_answer"),
difficulty=item.get("difficulty"),
)
return Question(
user_id=user_id,
type=parsed.type,
stem_markdown=parsed.stem_markdown,
options=[o.model_dump() for o in parsed.options] if parsed.options else None,
correct_answer=parsed.correct_answer,
difficulty=parsed.difficulty,
source_image_id=image_id,
ocr_engine=ocr.engine,
)
def _coerce_type(raw) -> QuestionType:
try:
return QuestionType(raw)
except ValueError:
return QuestionType.unclassified
def _fallback_question(
user_id: str, ocr, image_id: str | None, stem_override: str | None = None
) -> Question:
return Question(
user_id=user_id,
type=QuestionType.unclassified,
stem_markdown=stem_override or ocr.markdown,
source_image_id=image_id,
ocr_engine=ocr.engine,
)
async def _fail(db, job: ProcessingJob, msg: str):
job.status = JobStatus.failed
job.error_message = msg[:1000]
await db.commit()
+24
View File
@@ -0,0 +1,24 @@
"""客观题本地判分:精确比对答案集合。"""
from app.models.question import OBJECTIVE_TYPES, Question
def is_objective(q: Question) -> bool:
return q.type in OBJECTIVE_TYPES
def can_grade_locally(q: Question) -> bool:
"""客观题且录了标准答案,才能本地判。"""
return is_objective(q) and bool(q.correct_answer)
def grade_objective(q: Question, user_answer: list[str]) -> bool | None:
"""选择/判断题:忽略顺序与大小写,比对答案集合。
没录标准答案时返回 None(无法判定),而不是 False —— 否则一条只是
随手记下、还没填答案的笔记会被静默算错,并污染复习清单。
"""
correct = {str(a).strip().upper() for a in (q.correct_answer or [])}
if not correct:
return None
given = {str(a).strip().upper() for a in (user_answer or [])}
return correct == given
View File
+111
View File
@@ -0,0 +1,111 @@
"""AI 任务:判分、讲解、知识点总结。统一记录到 ai_generations。"""
import json
import logging
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.practice import AiGeneration
from app.models.question import Question
from app.models.user import UserSettings
from app.services.llm.client import LLMClient
from app.services.llm.prompts import (
EXPLAIN_SYSTEM,
EXPLAIN_USER_TEMPLATE,
JUDGE_SYSTEM,
JUDGE_USER_TEMPLATE,
TAGS_SYSTEM,
TAGS_USER_TEMPLATE,
)
logger = logging.getLogger(__name__)
async def load_llm(db: AsyncSession, user_id: str) -> LLMClient:
"""加载用户的 LLMClient(未配置会抛 LLMConfigError)。"""
row = await db.scalar(select(UserSettings).where(UserSettings.user_id == user_id))
return LLMClient(row)
def _parse_json(raw: str) -> dict:
text = raw.strip()
if text.startswith("```"):
text = text.split("```", 2)[1]
if text.startswith("json"):
text = text[4:]
return json.loads(text.strip())
async def _record(
db: AsyncSession, user_id: str, question_id: str | None, task: str, result, content: str
):
db.add(
AiGeneration(
user_id=user_id,
question_id=question_id,
task=task,
model=result.model,
request_tokens=result.prompt_tokens,
response_tokens=result.completion_tokens,
content_markdown=content,
)
)
async def judge_answer(
db: AsyncSession, llm: LLMClient, user_id: str, q: Question, user_answer: list[str]
) -> dict:
"""主观题判分。返回 {is_correct, feedback_markdown, key_points_missed}。"""
answer = "".join(q.correct_answer) if q.correct_answer else "(未提供标准答案)"
result = await llm.chat_text(
JUDGE_SYSTEM,
JUDGE_USER_TEMPLATE.format(
stem=q.stem_markdown, answer=answer, user_answer="\n".join(user_answer)
),
json_mode=True,
)
await _record(db, user_id, q.id, "judge", result, result.content)
try:
data = _parse_json(result.content)
return {
"is_correct": bool(data.get("is_correct")),
"feedback_markdown": data.get("feedback_markdown", ""),
"key_points_missed": data.get("key_points_missed", []),
}
except (json.JSONDecodeError, KeyError) as e:
logger.warning("judge 解析失败:%s", e)
return {
"is_correct": False,
"feedback_markdown": result.content,
"key_points_missed": [],
}
async def explain_question(
db: AsyncSession, llm: LLMClient, user_id: str, q: Question
) -> str:
answer_hint = ""
if q.correct_answer:
answer_hint = f"参考答案:{''.join(q.correct_answer)}\n\n"
result = await llm.chat_text(
EXPLAIN_SYSTEM,
EXPLAIN_USER_TEMPLATE.format(stem=q.stem_markdown, answer_hint=answer_hint),
)
await _record(db, user_id, q.id, "explain", result, result.content)
return result.content
async def summarize_tags(
db: AsyncSession, llm: LLMClient, user_id: str, q: Question
) -> list[str]:
result = await llm.chat_text(
TAGS_SYSTEM, TAGS_USER_TEMPLATE.format(stem=q.stem_markdown), json_mode=True
)
await _record(db, user_id, q.id, "summarize", result, result.content)
try:
data = _parse_json(result.content)
tags = data.get("tags", [])
return [str(t).strip() for t in tags if str(t).strip()][:4]
except (json.JSONDecodeError, KeyError) as e:
logger.warning("tags 解析失败:%s", e)
return []
+103
View File
@@ -0,0 +1,103 @@
"""统一的 OpenAI 兼容 LLM client:文本 + 多模态。
配置来自用户的 UserSettingsbase_url / api_key / 模型名)。api_key 在库里
是 Fernet 加密的,这里解密后使用。
"""
import base64
from dataclasses import dataclass
from openai import AsyncOpenAI
from app.models.user import UserSettings
from app.services.crypto import decrypt
class LLMConfigError(RuntimeError):
"""用户尚未正确配置 AI 接口。"""
@dataclass
class LLMResult:
content: str
model: str | None
prompt_tokens: int | None
completion_tokens: int | None
class LLMClient:
def __init__(self, settings: UserSettings):
if not settings or not settings.llm_base_url or not settings.llm_api_key_encrypted:
raise LLMConfigError("请先在设置里配置 AI 接口(base_url 与 api_key")
api_key = decrypt(settings.llm_api_key_encrypted)
if not api_key:
raise LLMConfigError("api_key 解密失败,请在设置里重新填写")
self._client = AsyncOpenAI(base_url=settings.llm_base_url, api_key=api_key)
self._text_model = settings.llm_text_model
self._vision_model = settings.llm_vision_model or settings.llm_text_model
async def chat_text(
self,
system: str,
user: str,
*,
json_mode: bool = False,
model: str | None = None,
) -> LLMResult:
m = model or self._text_model
if not m:
raise LLMConfigError("未配置文本模型名")
kwargs = {}
if json_mode:
kwargs["response_format"] = {"type": "json_object"}
resp = await self._client.chat.completions.create(
model=m,
messages=[
{"role": "system", "content": system},
{"role": "user", "content": user},
],
**kwargs,
)
return self._to_result(resp, m)
async def chat_vision(
self,
system: str,
user_text: str,
image_bytes: bytes,
mime: str = "image/jpeg",
*,
json_mode: bool = False,
model: str | None = None,
) -> LLMResult:
m = model or self._vision_model
if not m:
raise LLMConfigError("未配置多模态模型名")
data_url = f"data:{mime};base64,{base64.b64encode(image_bytes).decode()}"
kwargs = {}
if json_mode:
kwargs["response_format"] = {"type": "json_object"}
resp = await self._client.chat.completions.create(
model=m,
messages=[
{"role": "system", "content": system},
{
"role": "user",
"content": [
{"type": "text", "text": user_text},
{"type": "image_url", "image_url": {"url": data_url}},
],
},
],
**kwargs,
)
return self._to_result(resp, m)
@staticmethod
def _to_result(resp, model: str) -> LLMResult:
usage = getattr(resp, "usage", None)
return LLMResult(
content=resp.choices[0].message.content or "",
model=model,
prompt_tokens=getattr(usage, "prompt_tokens", None),
completion_tokens=getattr(usage, "completion_tokens", None),
)
+75
View File
@@ -0,0 +1,75 @@
"""各 AI 任务的 prompt 模板。"""
# ---- OCR:多模态直接看图,转成含 LaTeX 的 Markdown(不作答)----
VLM_OCR_SYSTEM = (
"你是一个精准的题目识别助手。用户会给你一张包含题目的图片,"
"请把图片中的题目内容原样转写为 Markdown 文本。要求:\n"
"1. 数学公式用 LaTeX 表示,行内用 $...$,独立公式用 $$...$$。\n"
"2. 保留题目的选项(如 A/B/C/D)、题号、结构。\n"
"3. 只转写题目本身,不要作答、不要解释、不要添加任何额外内容。\n"
"4. 如果有多道题,全部转写。"
)
VLM_OCR_USER = "请转写这张图片里的题目。"
# ---- 结构化提取:把 OCR/图片内容解析成结构化题目 JSON ----
EXTRACT_SYSTEM = """你是一个题目结构化助手。用户会给你一段题目文本(Markdown,可能含 LaTeX 公式),
其中可能包含一道或多道题。请把**每一道题**解析成一个对象,输出 JSON:
{
"questions": [
{
"type": "unclassified | single_choice | multiple_choice | true_false | fill_blank | short_answer",
"stem_markdown": "题干(保留原始 LaTeX",
"options": [{"key": "A", "text_markdown": "..."}] // 仅选择/判断题;否则为 null,
"correct_answer": ["A"] 或 ["文本答案"] 或 null, // 原文未给出答案则为 null
"difficulty": null 或 1-5
}
]
}
题型判断规则:
- 有 ABCD 等多个选项且只选一个 → single_choice;可多选 → multiple_choice
- 判断对错(正确/错误、对/错、True/False)→ true_falseoptions 用 [{"key":"T","text_markdown":"正确"},{"key":"F","text_markdown":"错误"}]
- 有下划线、括号、填空 → fill_blank
- 其余开放性问答 → short_answer
- 判断不了、或者这段内容不像一道完整的题 → unclassified(这是可以接受的,不要硬猜)
重要:
- 有几道题就输出几个对象,不要合并成一条,也不要漏掉。
- 只有一道题时也要放在 questions 数组里。
- 原文没给答案就填 null,不要自己解题填答案。
只输出 JSON,不要额外文字。"""
EXTRACT_USER_TEMPLATE = "题目文本:\n\n{content}"
# ---- 判分(主观题)----
JUDGE_SYSTEM = """你是一个阅卷助手。根据题目、标准答案(可能没有)和学生答案,判断对错并给出反馈。
输出 JSON
{
"is_correct": true/false,
"feedback_markdown": "简短反馈,指出对错原因(可含 LaTeX)",
"key_points_missed": ["遗漏的要点", ...]
}
只输出 JSON。"""
JUDGE_USER_TEMPLATE = """题目:
{stem}
标准答案:{answer}
学生答案:
{user_answer}"""
# ---- 讲解 ----
EXPLAIN_SYSTEM = (
"你是一位耐心的老师。请针对给定题目给出面向自学者的分步讲解,"
"解释解题思路而非只给结论。使用 Markdown,公式用 LaTeX$...$ / $$...$$)。"
)
EXPLAIN_USER_TEMPLATE = """题目:
{stem}
{answer_hint}请给出详细讲解。"""
# ---- 知识点标签 ----
TAGS_SYSTEM = """你是一个知识点归类助手。根据题目内容,输出 1-4 个简短的知识点标签。
输出 JSON{"tags": ["标签1", "标签2"]}
标签要简洁(2-8 字),是学科知识点,不要句子。只输出 JSON。"""
TAGS_USER_TEMPLATE = "题目:\n{stem}"
View File
+14
View File
@@ -0,0 +1,14 @@
"""OCR 引擎抽象接口。"""
from dataclasses import dataclass
from typing import Protocol
@dataclass
class OcrResult:
markdown: str
confidence: float # 0-1,未知时给 1.0
engine: str # "pix2text" / "vlm"
class OcrEngine(Protocol):
async def recognize(self, image_bytes: bytes, mime: str) -> OcrResult: ...
+35
View File
@@ -0,0 +1,35 @@
"""OCR 编排:Pix2Text 优先,失败/低置信/未部署时降级到多模态 VLM。"""
import logging
from app.config import settings
from app.services.llm.client import LLMClient
from app.services.ocr.base import OcrResult
from app.services.ocr.pix2text_client import Pix2TextClient
from app.services.ocr.vlm_engine import VlmOcrEngine
logger = logging.getLogger(__name__)
# Pix2Text 结果置信度低于此值则降级到 VLM
_CONFIDENCE_THRESHOLD = 0.5
async def run_ocr(image_bytes: bytes, mime: str, llm: LLMClient) -> OcrResult:
"""返回识别结果,engine 字段标明最终使用的引擎。"""
# 1. 若配置了 Pix2Text 服务且健康,优先用
if settings.ocr_service_url:
client = Pix2TextClient(settings.ocr_service_url)
if await client.healthy():
try:
result = await client.recognize(image_bytes, mime)
if result.markdown.strip() and result.confidence >= _CONFIDENCE_THRESHOLD:
return result
logger.info(
"Pix2Text 置信度低(%.2f)或结果为空,降级 VLM", result.confidence
)
except Exception as e: # noqa: BLE001
logger.warning("Pix2Text 识别失败,降级 VLM%s", e)
else:
logger.info("Pix2Text 服务不健康,降级 VLM")
# 2. 降级 / 默认:多模态大模型
return await VlmOcrEngine(llm).recognize(image_bytes, mime)
+32
View File
@@ -0,0 +1,32 @@
"""Pix2Text 独立微服务的 HTTP 客户端(后续阶段部署 OCR 服务时启用)。"""
import httpx
from app.services.ocr.base import OcrResult
class Pix2TextClient:
def __init__(self, service_url: str, timeout: float = 60.0):
self._url = service_url.rstrip("/")
self._timeout = timeout
async def healthy(self) -> bool:
try:
async with httpx.AsyncClient(timeout=5.0) as c:
r = await c.get(f"{self._url}/health")
return r.status_code == 200
except httpx.HTTPError:
return False
async def recognize(self, image_bytes: bytes, mime: str) -> OcrResult:
async with httpx.AsyncClient(timeout=self._timeout) as c:
r = await c.post(
f"{self._url}/ocr",
files={"file": ("image", image_bytes, mime)},
)
r.raise_for_status()
data = r.json()
return OcrResult(
markdown=data["markdown"],
confidence=float(data.get("confidence", 1.0)),
engine="pix2text",
)
+15
View File
@@ -0,0 +1,15 @@
"""多模态大模型 OCR 引擎:直接看图转写题目为 Markdown。"""
from app.services.llm.client import LLMClient
from app.services.llm.prompts import VLM_OCR_SYSTEM, VLM_OCR_USER
from app.services.ocr.base import OcrResult
class VlmOcrEngine:
def __init__(self, llm: LLMClient):
self._llm = llm
async def recognize(self, image_bytes: bytes, mime: str) -> OcrResult:
result = await self._llm.chat_vision(
VLM_OCR_SYSTEM, VLM_OCR_USER, image_bytes, mime=mime
)
return OcrResult(markdown=result.content, confidence=1.0, engine="vlm")
+43
View File
@@ -0,0 +1,43 @@
"""密码哈希(bcrypt)与 JWT 签发/校验。"""
from datetime import datetime, timedelta, timezone
import bcrypt
import jwt
from app.config import settings
_BCRYPT_ROUNDS = 12
def hash_password(password: str) -> str:
hashed = bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt(rounds=_BCRYPT_ROUNDS))
return hashed.decode("utf-8")
def verify_password(password: str, password_hash: str) -> bool:
try:
return bcrypt.checkpw(password.encode("utf-8"), password_hash.encode("utf-8"))
except ValueError:
return False
def create_access_token(subject: str) -> str:
"""subject 为 user id。"""
now = datetime.now(timezone.utc)
payload = {
"sub": subject,
"iat": now,
"exp": now + timedelta(minutes=settings.access_token_expire_minutes),
}
return jwt.encode(payload, settings.jwt_secret, algorithm=settings.jwt_algorithm)
def decode_access_token(token: str) -> str | None:
"""返回 user id,失败返回 None。"""
try:
payload = jwt.decode(
token, settings.jwt_secret, algorithms=[settings.jwt_algorithm]
)
return payload.get("sub")
except jwt.PyJWTError:
return None
+38
View File
@@ -0,0 +1,38 @@
"""启动时的清理工作。"""
import logging
from sqlalchemy import update
from app.db import SessionLocal
from app.models.job import JobStatus, ProcessingJob
logger = logging.getLogger(__name__)
# 进程重启后不可能再继续的中间状态
_STALE_STATUSES = (
JobStatus.pending,
JobStatus.ocr_running,
JobStatus.ai_running,
)
async def fail_stale_jobs() -> int:
"""把残留的进行中任务标记为失败。
任务跑在 BackgroundTasks 里,进程重启就没了。不清理的话这些任务会永远
停在 pending/*_running,前端会一直轮询下去。
"""
async with SessionLocal() as db:
result = await db.execute(
update(ProcessingJob)
.where(ProcessingJob.status.in_(_STALE_STATUSES))
.values(
status=JobStatus.failed,
error_message="服务重启,任务中断,请重新上传",
)
)
await db.commit()
count = result.rowcount or 0
if count:
logger.info("已清理 %d 个残留任务", count)
return count
+99
View File
@@ -0,0 +1,99 @@
"""对象存储(S3 兼容):上传/下载图片,含校验与去重。
boto3 是同步库,直接在 async 路径里调用会阻塞事件循环(传大图时整个服务卡住),
所以对外暴露 async 版本(to_thread 包装),同步版本保留给测试和脚本用。
"""
import asyncio
import hashlib
import io
import uuid
import boto3
from botocore.config import Config
from PIL import Image as PILImage
from app.config import settings
# Pillow 格式 → 扩展名
_FORMAT_EXT = {"JPEG": "jpg", "PNG": "png", "WEBP": "webp"}
# 解压炸弹防护:单边像素上限
_MAX_DIMENSION = 12000
class StorageError(RuntimeError):
pass
class ImageValidationError(ValueError):
pass
def _s3_client():
if not settings.s3_bucket:
raise StorageError("S3 未配置(S3_BUCKET 为空)")
return boto3.client(
"s3",
endpoint_url=settings.s3_endpoint_url,
region_name=settings.s3_region,
aws_access_key_id=settings.s3_access_key,
aws_secret_access_key=settings.s3_secret_key,
config=Config(signature_version="s3v4"),
)
def validate_image(data: bytes) -> tuple[str, str]:
"""校验是真实图片且格式在白名单内。返回 (mime, ext)。"""
if len(data) > settings.max_upload_bytes:
raise ImageValidationError("文件超过大小上限")
try:
img = PILImage.open(io.BytesIO(data))
img.verify() # 校验完整性(魔数级)
# verify 后需重新打开才能读属性
img = PILImage.open(io.BytesIO(data))
except Exception as e: # noqa: BLE001
raise ImageValidationError(f"不是有效的图片文件:{e}") from e
fmt = img.format or ""
if fmt not in _FORMAT_EXT:
raise ImageValidationError(f"不支持的图片格式:{fmt}")
if img.width > _MAX_DIMENSION or img.height > _MAX_DIMENSION:
raise ImageValidationError("图片尺寸过大")
mime = f"image/{'jpeg' if fmt == 'JPEG' else fmt.lower()}"
if mime not in settings.allowed_image_mimes:
raise ImageValidationError(f"不支持的 MIME{mime}")
return mime, _FORMAT_EXT[fmt]
def upload_image(user_id: str, data: bytes, ext: str, mime: str) -> str:
"""上传到 S3,返回 object_key。"""
key = f"{user_id}/{uuid.uuid4().hex}.{ext}"
client = _s3_client()
client.put_object(
Bucket=settings.s3_bucket,
Key=key,
Body=data,
ContentType=mime,
)
return key
def download_image(object_key: str) -> bytes:
client = _s3_client()
resp = client.get_object(Bucket=settings.s3_bucket, Key=object_key)
return resp["Body"].read()
def sha256_hex(data: bytes) -> str:
return hashlib.sha256(data).hexdigest()
# ---- async 包装:在 async 路径里请用这两个 ----
async def upload_image_async(user_id: str, data: bytes, ext: str, mime: str) -> str:
return await asyncio.to_thread(upload_image, user_id, data, ext, mime)
async def download_image_async(object_key: str) -> bytes:
return await asyncio.to_thread(download_image, object_key)
+36
View File
@@ -0,0 +1,36 @@
"""标签 get-or-create(按用户隔离)。"""
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.question import Tag
async def get_or_create_tags(
db: AsyncSession, user_id: str, names: list[str]
) -> list[Tag]:
"""返回该用户下对应名字的 Tag,不存在则创建。不提交(由调用方 commit)。"""
cleaned = [n.strip() for n in names if n and n.strip()]
if not cleaned:
return []
# 去重保序
seen: dict[str, None] = {}
for n in cleaned:
seen.setdefault(n, None)
unique = list(seen.keys())
existing = (
await db.scalars(
select(Tag).where(Tag.user_id == user_id, Tag.name.in_(unique))
)
).all()
by_name = {t.name: t for t in existing}
result: list[Tag] = []
for name in unique:
tag = by_name.get(name)
if tag is None:
tag = Tag(user_id=user_id, name=name)
db.add(tag)
by_name[name] = tag
result.append(tag)
return result