feat:增加知识库数据接口

This commit is contained in:
2026-07-28 19:54:50 +08:00
parent 2e6e571973
commit 0e9bde9a14
19 changed files with 703 additions and 31 deletions

8
.env
View File

@@ -7,6 +7,8 @@ OCR_API_URL=https://paddleocr.aistudio-app.com/api/v2/ocr/jobs
OCR_API_TOKEN=df7dcc85a5c3c9d64e421f353d11d13ec45512f6
OCR_MODEL=PaddleOCR-VL-1.6
MODEL_NAME=deepseek-v4-flash
MODEL_BASE_URL=https://api.deepseek.com
MODEL_API_KEY=sk-0b237d41f6bc44fc9732ea66bd7eade0
MODEL_NAME="glm-5.2"
MODEL_BASE_URL="https://dashscope.aliyuncs.com/compatible-mode/v1"
MODEL_API_KEY="sk-52bcd98e9c1d45908437c4e8706eefff"
DATABASE_URL=postgresql+asyncpg://agent:medical%402019@60.247.145.200:5432/study

View File

@@ -11,3 +11,7 @@ model = init_chat_model(
base_url=os.getenv('MODEL_BASE_URL'),
api_key=os.getenv('MODEL_API_KEY')
)
# MODEL_NAME=deepseek-v4-flash
# MODEL_BASE_URL=https://api.deepseek.com
# MODEL_API_KEY=sk-0b237d41f6bc44fc9732ea66bd7eade0

View File

@@ -1,27 +1,50 @@
study_prompt = """
你是一位专业的出题专家,擅长根据学习材料生成高质量选择题
你是一位专业的出题专家,擅长根据学习材料生成高质量的客观题(单选题和多选题)
要求:
1. 题目必须严格基于学习内容,不得编造知识点。
2. 每道题只有 1 个正确答案。
3. 干扰项要有迷惑性。
4. 输出必须是纯 JSON禁止 Markdown。
5. 不要修改专业术语。
请根据下方【学习材料】生成 510 道题目,严格遵守以下要求:
返回格式:
{
"questions": [
1. **题型要求**
- 可包含单选题和多选题。
- 多选题至少 2 个正确选项,最多不超过 4 个。
- 每道题的选项固定为 4 项A/B/C/D
2. **出题原则**
- 所有题目必须严格基于【学习材料】,不得编造知识点。
- 干扰项应具有一定迷惑性,但不能是完全无关或明显错误的内容。
- 专业术语必须与材料保持一致,禁止改写、简写或口语化。
- 题目表述必须自然、规范,禁止出现“根据材料”“根据以上内容”“原文指出”等提示语。
- 题目应像正式考试题一样直接提问。
3. **输出要求(非常重要)**
- 仅返回纯 JSON禁止输出 Markdown、代码块、注释或任何说明文字。
- 必须返回一个 JSON 数组,不得以对象包裹。
- JSON 必须是合法格式,可被程序直接解析。
4. **返回格式(严格遵循)**
[
{
"id": 1,
"type": "single",
"question": "题目内容",
"options": {
"A": "选项A内容",
"B": "选项B内容",
"C": "选项C内容",
"D": "选项D内容"
"options": [
"选项A内容",
"选项B内容",
"选项C内容",
"选项D内容"
],
"answer": ["A"]
},
"answer": "A"
{
"id": 2,
"type": "multiple",
"question": "题目内容",
"options": [
"选项A内容",
"选项B内容",
"选项C内容",
"选项D内容"
],
"answer": ["A", "C"]
}
]
}
]
"""

27
database.py Normal file
View File

@@ -0,0 +1,27 @@
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
from sqlalchemy.orm import declarative_base, sessionmaker
from dotenv import load_dotenv
import os
load_dotenv()
DATABASE_URL = os.getenv("DATABASE_URL")
engine = create_async_engine(
DATABASE_URL,
echo=True,
future=True
)
AsyncSessionLocal = sessionmaker(
bind=engine,
class_=AsyncSession,
expire_on_commit=False
)
Base = declarative_base()
async def get_db():
async with AsyncSessionLocal() as session:
yield session

0
id_generator/__init__.py Normal file
View File

38
id_generator/generator.py Normal file
View File

@@ -0,0 +1,38 @@
"""
雪花算法生成器IdGenerator
"""
# !/usr/bin/python
# coding=UTF-8
from . import options
from . import snowflake_m1
class DefaultIdGenerator:
"""
ID生成器
"""
def __init__(self):
self.snowflake = None
def set_id_generator(self, option: options.IdGeneratorOptions):
"""
设置id生成规则信息
"""
if option.base_time < 100000:
raise ValueError("base time error.")
self.snowflake = snowflake_m1.SnowFlakeM1(option)
def next_id(self) -> int:
"""
获取新的UUID
"""
if self.snowflake is None:
raise ValueError("please set id generator at first.")
return self.snowflake.next_id()

134
id_generator/idregister.py Normal file
View File

@@ -0,0 +1,134 @@
"""
worker id generator
"""
# !/usr/bin/python
# coding=UTF-8
from threading import Thread
import time
import logging
import redis
class Register:
"""
redis封装
- host 代表redis ip
- port 代表redis端口
- max_worker_id worker_id的最大值, 默认为100
- password redis的密码, 默认为空
"""
def __init__(self, host, port, max_worker_id=100, password=None):
self.redis_impl = redis.StrictRedis(host=host, port=port, db=0, password=password)
self.loop_count = 0
self.max_loop_count = 10
self.worker_id_expire_time = 15
self.max_worker_id = max_worker_id
self.worker_id = -1
self.is_stop = False
def get_lock(self, key):
"""
获取分布式全局锁,并设置过期时间为30秒
"""
if self.redis_impl.setnx(key, 1):
self.redis_impl.expire(key, 30)
return True
if self.redis_impl.ttl(key) < 0:
self.redis_impl.expire(key, 30)
return False
def stop(self):
"""
退出注册器的线程
"""
self.is_stop = True
def get_worker_id(self):
"""
获取全局唯一worker_id, 会创建一个线程给worker id续期
失败返回-1
"""
self.loop_count = 0
def extern_life(my_id):
while 1:
time.sleep(self.worker_id_expire_time / 3)
# 是否关闭了
if self.is_stop:
return
# 更新生命周期
if self.worker_id != my_id:
break
try:
self.redis_impl.expire(
f"IdGen:WorkerId:Value:{my_id}",
self.worker_id_expire_time)
except Exception as exe:
logging.error(exe)
continue
self.worker_id = self.__get_next_worker_id()
if self.worker_id > -1:
Thread(target=extern_life, args=[self.worker_id]).start()
return self.worker_id
def __get_next_worker_id(self):
"""
获取全局唯一worker id内部实现
"""
cur = self.redis_impl.incrby("IdGen:WorkerId:Index", 1)
def can_reset():
try:
reset_value = self.redis_impl.incr("IdGen:WorkerId:Value:Edit")
return reset_value != 1
except Exception as ept:
logging.error(ept)
return False
def end_reset():
try:
self.redis_impl.set("IdGen:WorkerId:Value:Edit", 0)
except Exception as ept:
logging.error(ept)
def is_available(worker_id: int):
try:
rst = self.redis_impl.get(f"IdGen:WorkerId:Value:{worker_id}")
return rst != "Y"
except Exception as ept:
logging.error(ept)
return False
if cur > self.max_worker_id:
if can_reset():
self.redis_impl.set("IdGen:WorkerId:Index", -1)
end_reset()
self.loop_count += 1
if self.loop_count > self.max_loop_count:
self.loop_count = 0
return -1
time.sleep(0.2 * self.loop_count)
return self.__get_next_worker_id()
time.sleep(0.2)
return self.__get_next_worker_id()
if is_available(cur):
self.redis_impl.setex(
f"IdGen:WorkerId:Value:{cur}",
self.worker_id_expire_time,
"Y"
)
self.loop_count = 0
return cur
return self.__get_next_worker_id()

43
id_generator/options.py Normal file
View File

@@ -0,0 +1,43 @@
"""
生成器IdGenerator配置选项
"""
# !/usr/bin/python
# coding=UTF-8
class IdGeneratorOptions:
"""
ID生成器配置
- worker_id 全局唯一id, 区分不同uuid生成器实例
- worker_id_bit_length 生成的uuid中worker_id占用的位数
- seq_bit_length 生成的uuid中序列号占用的位数
"""
def __init__(self, worker_id=0, worker_id_bit_length=6, seq_bit_length=6):
# 雪花计算方法,1-漂移算法|2-传统算法), 默认1。目前只实现了1。
self.method = 1
# 基础时间ms单位, 不能超过当前系统时间
self.base_time = 1582136402000
# 机器码, 必须由外部设定, 最大值 2^worker_id_bit_length-1
self.worker_id = worker_id
# 机器码位长, 默认值6, 取值范围 [1, 15](要求:序列数位长+机器码位长不超过22
self.worker_id_bit_length = worker_id_bit_length
# 序列数位长, 默认值6, 取值范围 [3, 21](要求:序列数位长+机器码位长不超过22
self.seq_bit_length = seq_bit_length
# 最大序列数(含), 设置范围 [max_seq_number, 2^seq_bit_length-1]
# 默认值0, 表示最大序列数取最大值2^seq_bit_length-1]
self.max_seq_number = 0
# 最小序列数(含), 默认值5, 取值范围 [5, max_seq_number], 每毫秒的前5个序列数对应编号0-4是保留位
# 其中1-4是时间回拨相应预留位, 0是手工新值预留位
self.min_seq_number = 5
# 最大漂移次数(含), 默认2000, 推荐范围500-10000与计算能力有关
self.top_over_cost_count = 2000

20
id_generator/snowflake.py Normal file
View File

@@ -0,0 +1,20 @@
"""
雪花算法生成器接口声明
"""
# !/usr/bin/python
# coding=UTF-8
class SnowFlake():
def __init__(self, options):
self.options = options
def next_id(self) -> int:
"""
获取新的UUID
"""
return 0

View File

@@ -0,0 +1,147 @@
"""
M1生成器
"""
# !/usr/bin/python
# coding=UTF-8
import threading
import time
from .snowflake import SnowFlake
from .options import IdGeneratorOptions
class SnowFlakeM1(SnowFlake):
"""
M1规则ID生成器配置
"""
def __init__(self, options: IdGeneratorOptions):
# 1.base_time
self.base_time = 1582136402000
if options.base_time != 0:
self.base_time = int(options.base_time)
# 2.worker_id_bit_length
self.worker_id_bit_length = 6
if options.worker_id_bit_length != 0:
self.worker_id_bit_length = int(options.worker_id_bit_length)
# 3.worker_id
self.worker_id = options.worker_id
# 4.seq_bit_length
self.seq_bit_length = 6
if options.seq_bit_length != 0:
self.seq_bit_length = int(options.seq_bit_length)
# 5.max_seq_number
self.max_seq_number = int(options.max_seq_number)
if options.max_seq_number <= 0:
self.max_seq_number = (1 << self.seq_bit_length) - 1
# 6.min_seq_number
self.min_seq_number = int(options.min_seq_number)
# 7.top_over_cost_count
self.top_over_cost_count = int(options.top_over_cost_count)
# 8.Others
self.__timestamp_shift = self.worker_id_bit_length + self.seq_bit_length
self.__current_seq_number = self.min_seq_number
self.__last_time_tick: int = 0
self.__turn_back_time_tick: int = 0
self.__turn_back_index: int = 0
self.__is_over_cost = False
self.___over_cost_count_in_one_term: int = 0
self.__id_lock = threading.Lock()
def __next_over_cost_id(self) -> int:
current_time_tick = self.__get_current_time_tick()
if current_time_tick > self.__last_time_tick:
self.__last_time_tick = current_time_tick
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = False
self.___over_cost_count_in_one_term = 0
return self.__calc_id(self.__last_time_tick)
if self.___over_cost_count_in_one_term >= self.top_over_cost_count:
self.__last_time_tick = self.__get_next_time_tick()
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = False
self.___over_cost_count_in_one_term = 0
return self.__calc_id(self.__last_time_tick)
if self.__current_seq_number > self.max_seq_number:
self.__last_time_tick += 1
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = True
self.___over_cost_count_in_one_term += 1
return self.__calc_id(self.__last_time_tick)
return self.__calc_id(self.__last_time_tick)
def __next_normal_id(self) -> int:
current_time_tick = self.__get_current_time_tick()
if current_time_tick < self.__last_time_tick:
if self.__turn_back_time_tick < 1:
self.__turn_back_time_tick = self.__last_time_tick - 1
self.__turn_back_index += 1
# 每毫秒序列数的前5位是预留位, 0用于手工新值, 1-4是时间回拨次序
# 支持4次回拨次序避免回拨重叠导致ID重复, 可无限次回拨(次序循环使用)。
if self.__turn_back_index > 4:
self.__turn_back_index = 1
return self.__calc_turn_back_id(self.__turn_back_time_tick)
# 时间追平时, _TurnBackTimeTick清零
self.__turn_back_time_tick = min(self.__turn_back_time_tick, 0)
if current_time_tick > self.__last_time_tick:
self.__last_time_tick = current_time_tick
self.__current_seq_number = self.min_seq_number
return self.__calc_id(self.__last_time_tick)
if self.__current_seq_number > self.max_seq_number:
self.__last_time_tick += 1
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = True
self.___over_cost_count_in_one_term = 1
return self.__calc_id(self.__last_time_tick)
return self.__calc_id(self.__last_time_tick)
def __calc_id(self, use_time_tick) -> int:
self.__current_seq_number += 1
return (
(use_time_tick << self.__timestamp_shift) +
(self.worker_id << self.seq_bit_length) +
self.__current_seq_number
) % int(1e64)
def __calc_turn_back_id(self, use_time_tick) -> int:
self.__turn_back_time_tick -= 1
return (
(use_time_tick << self.__timestamp_shift) +
(self.worker_id << self.seq_bit_length) +
self.__turn_back_index
) % int(1e64)
def __get_current_time_tick(self) -> int:
return int((time.time_ns() / 1e6) - self.base_time)
def __get_next_time_tick(self) -> int:
temp_time_ticker = self.__get_current_time_tick()
while temp_time_ticker <= self.__last_time_tick:
# 0.001 = 1 mili sec
time.sleep(0.001)
temp_time_ticker = self.__get_current_time_tick()
return temp_time_ticker
def next_id(self) -> int:
with self.__id_lock:
if self.__is_over_cost:
nextid = self.__next_over_cost_id()
else:
nextid = self.__next_normal_id()
return nextid

28
main.py
View File

@@ -1,9 +1,15 @@
from fastapi import FastAPI, UploadFile, File
from typing import List
from fastapi import FastAPI, UploadFile, File, Depends
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.middleware.cors import CORSMiddleware
from starlette.responses import StreamingResponse
from models.agent import QueryRequest
from service import query_agent
from database import get_db
from schemas.agent import QueryRequest
from schemas.library import LibraryFileRequest, LibraryFileResponse
from service.agent_service import query_agent
from service.library_service import add_file, get_file, list_files
from storage import upload_rustfs
from ocr import ocr_from_url
@@ -34,3 +40,19 @@ def generate(query: QueryRequest):
query_agent(query),
media_type="text/event-stream"
)
@app.post("/files", response_model=bool)
async def add_library_file(data: LibraryFileRequest, db: AsyncSession = Depends(get_db)):
await add_file(db, data)
return True
@app.get("/files", response_model=List[LibraryFileResponse])
async def list_library_file(db: AsyncSession = Depends(get_db)):
return await list_files(db)
@app.get("/files/{file_id}", response_model=LibraryFileResponse)
async def get_library_file(file_id: int, db: AsyncSession = Depends(get_db)):
return await get_file(db, file_id)

43
models/base.py Normal file
View File

@@ -0,0 +1,43 @@
from datetime import datetime
from sqlalchemy import Column, BigInteger, String, event, TIMESTAMP, func
from sqlalchemy.ext.declarative import declared_attr
from utils.common import camel_to_snake
from database import Base
from id_generator import options, generator
# https://github.com/yitter/IdGenerator/tree/master/Python
options = options.IdGeneratorOptions(worker_id=23)
idgen = generator.DefaultIdGenerator()
idgen.set_id_generator(options)
# 第一层基类包含ID
class IdBase(Base):
__abstract__ = True
id = Column(BigInteger, primary_key=True, index=True)
@declared_attr
def __tablename__(cls):
# 自动把数据库实体类名驼峰转为数据库表名下划线
return camel_to_snake(cls.__name__)
# 自动填充id
@event.listens_for(IdBase, 'before_insert', propagate=True)
def before_insert_listener(mapper, connection, target):
if target.id is None:
target.id = idgen.next_id()
# 第二层基类包含ID和审计字段
class AuditBase(IdBase):
__abstract__ = True
create_time = Column(TIMESTAMP, nullable=False, default=datetime.now)
create_by = Column(String(255), nullable=True)
update_time = Column(TIMESTAMP, nullable=False, default=datetime.now, onupdate=datetime.now)
update_by = Column(String(255), nullable=True)

33
models/library.py Normal file
View File

@@ -0,0 +1,33 @@
from sqlalchemy import Column, BigInteger, String, Text, TIMESTAMP, Integer, ForeignKey
from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import relationship
from models.base import AuditBase
class LibraryFile(AuditBase):
subject = Column(String(45), nullable=False)
module = Column(String(45), nullable=False)
name = Column(String(45), nullable=False)
type = Column(String(45), nullable=False)
size = Column(Integer, nullable=False)
questions = relationship(
"LibraryFileQuestion",
back_populates="library_file",
cascade="all, delete-orphan"
)
class LibraryFileQuestion(AuditBase):
library_file_id = Column(BigInteger, ForeignKey("library_file.id"), nullable=False)
type = Column(String(45), nullable=False)
question = Column(Text, nullable=False)
options = Column(JSONB, nullable=False)
answer = Column(JSONB, nullable=False)
library_file = relationship(
"LibraryFile",
back_populates="questions"
)

View File

@@ -8,3 +8,4 @@ langchain-core~=1.5.1
langchain-openai~=1.4.1
starlette~=1.3.1
pydantic~=2.13.4
SQLAlchemy~=2.0.51

37
schemas/library.py Normal file
View File

@@ -0,0 +1,37 @@
from datetime import datetime
from pydantic import BaseModel
from typing import List
class QuestionRequest(BaseModel):
type: str
question: str
options: List[str]
answer: List[str]
class QuestionResponse(QuestionRequest):
id: int
class Config:
from_attributes = True
class LibraryFileRequest(BaseModel):
subject: str
module: str
name: str
size: int
type: str
questions: List[QuestionRequest] = []
class LibraryFileResponse(LibraryFileRequest):
id: int
uploadTime: str
questions: List[QuestionResponse] = []
class Config:
from_attributes = True

View File

@@ -1,7 +1,7 @@
from langchain_core.messages import AIMessageChunk, HumanMessage
from agent.study import study_agent
from models.agent import QueryRequest
from schemas.agent import QueryRequest
def query_agent(query: QueryRequest):

View File

@@ -0,0 +1,85 @@
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from models.library import LibraryFile, LibraryFileQuestion
from schemas.library import LibraryFileRequest, LibraryFileResponse, QuestionResponse
async def add_file(db: AsyncSession, data: LibraryFileRequest) -> bool:
"""
创建题库文件及其题目
"""
file = LibraryFile(
subject=data.subject,
module=data.module,
name=data.name,
size=data.size,
type=data.type
)
db.add(file)
await db.flush()
for q in data.questions:
question = LibraryFileQuestion(
library_file_id=file.id,
type=q.type,
question=q.question,
options=q.options,
answer=q.answer
)
db.add(question)
await db.commit()
return True
async def get_file(db: AsyncSession, file_id: int) -> LibraryFileResponse:
stmt = (
select(LibraryFile)
.options(selectinload(LibraryFile.questions))
.where(LibraryFile.id == file_id)
)
result = await db.execute(stmt)
file = result.scalar_one_or_none()
return LibraryFileResponse(
id=file.id,
subject=file.subject,
module=file.module,
name=file.name,
type=file.type,
size=file.size,
uploadTime=file.create_time.strftime("%Y-%m-%d %H:%M:%S"),
questions=[
QuestionResponse(
id=q.id,
type=q.type,
question=q.question,
options=q.options,
answer=q.answer,
)
for q in file.questions
],
)
async def list_files(db: AsyncSession) -> list[LibraryFileResponse]:
stmt = select(LibraryFile).order_by(LibraryFile.create_time.desc())
files = (await db.execute(stmt)).scalars().all()
return [
LibraryFileResponse(
id=f.id,
subject=f.subject,
module=f.module,
name=f.name,
type=f.type,
size=f.size,
uploadTime=f.create_time.strftime("%Y-%m-%d %H:%M:%S"),
questions=[],
)
for f in files
]

13
utils/common.py Normal file
View File

@@ -0,0 +1,13 @@
import re
def camel_to_snake(name: str) -> str:
"""将驼峰命名转换为蛇形命名CamelCase → snake_case"""
name = re.sub('(.)([A-Z][a-z]+)', r'\1_\2', name)
return re.sub('([a-z0-9])([A-Z])', r'\1_\2', name).lower()
def snake_to_camel(name: str) -> str:
"""将蛇形命名转换为驼峰命名snake_case → CamelCase"""
components = name.split('_')
return ''.join(x.title() for x in components)