feat:增加知识库数据接口
This commit is contained in:
8
.env
8
.env
@@ -7,6 +7,8 @@ OCR_API_URL=https://paddleocr.aistudio-app.com/api/v2/ocr/jobs
|
|||||||
OCR_API_TOKEN=df7dcc85a5c3c9d64e421f353d11d13ec45512f6
|
OCR_API_TOKEN=df7dcc85a5c3c9d64e421f353d11d13ec45512f6
|
||||||
OCR_MODEL=PaddleOCR-VL-1.6
|
OCR_MODEL=PaddleOCR-VL-1.6
|
||||||
|
|
||||||
MODEL_NAME=deepseek-v4-flash
|
MODEL_NAME="glm-5.2"
|
||||||
MODEL_BASE_URL=https://api.deepseek.com
|
MODEL_BASE_URL="https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||||
MODEL_API_KEY=sk-0b237d41f6bc44fc9732ea66bd7eade0
|
MODEL_API_KEY="sk-52bcd98e9c1d45908437c4e8706eefff"
|
||||||
|
|
||||||
|
DATABASE_URL=postgresql+asyncpg://agent:medical%402019@60.247.145.200:5432/study
|
||||||
|
|||||||
@@ -11,3 +11,7 @@ model = init_chat_model(
|
|||||||
base_url=os.getenv('MODEL_BASE_URL'),
|
base_url=os.getenv('MODEL_BASE_URL'),
|
||||||
api_key=os.getenv('MODEL_API_KEY')
|
api_key=os.getenv('MODEL_API_KEY')
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# MODEL_NAME=deepseek-v4-flash
|
||||||
|
# MODEL_BASE_URL=https://api.deepseek.com
|
||||||
|
# MODEL_API_KEY=sk-0b237d41f6bc44fc9732ea66bd7eade0
|
||||||
|
|||||||
@@ -1,27 +1,50 @@
|
|||||||
study_prompt = """
|
study_prompt = """
|
||||||
你是一位专业的出题专家,擅长根据学习材料生成高质量选择题。
|
你是一位专业的出题专家,擅长根据学习材料生成高质量的客观题(单选题和多选题)。
|
||||||
|
|
||||||
要求:
|
请根据下方【学习材料】生成 5~10 道题目,严格遵守以下要求:
|
||||||
1. 题目必须严格基于学习内容,不得编造知识点。
|
|
||||||
2. 每道题只有 1 个正确答案。
|
|
||||||
3. 干扰项要有迷惑性。
|
|
||||||
4. 输出必须是纯 JSON,禁止 Markdown。
|
|
||||||
5. 不要修改专业术语。
|
|
||||||
|
|
||||||
返回格式:
|
1. **题型要求**
|
||||||
{
|
- 可包含单选题和多选题。
|
||||||
"questions": [
|
- 多选题至少 2 个正确选项,最多不超过 4 个。
|
||||||
|
- 每道题的选项固定为 4 项(A/B/C/D)。
|
||||||
|
|
||||||
|
2. **出题原则**
|
||||||
|
- 所有题目必须严格基于【学习材料】,不得编造知识点。
|
||||||
|
- 干扰项应具有一定迷惑性,但不能是完全无关或明显错误的内容。
|
||||||
|
- 专业术语必须与材料保持一致,禁止改写、简写或口语化。
|
||||||
|
- 题目表述必须自然、规范,禁止出现“根据材料”“根据以上内容”“原文指出”等提示语。
|
||||||
|
- 题目应像正式考试题一样直接提问。
|
||||||
|
|
||||||
|
3. **输出要求(非常重要)**
|
||||||
|
- 仅返回纯 JSON,禁止输出 Markdown、代码块、注释或任何说明文字。
|
||||||
|
- 必须返回一个 JSON 数组,不得以对象包裹。
|
||||||
|
- JSON 必须是合法格式,可被程序直接解析。
|
||||||
|
|
||||||
|
4. **返回格式(严格遵循)**
|
||||||
|
[
|
||||||
{
|
{
|
||||||
"id": 1,
|
"id": 1,
|
||||||
|
"type": "single",
|
||||||
"question": "题目内容",
|
"question": "题目内容",
|
||||||
"options": {
|
"options": [
|
||||||
"A": "选项A内容",
|
"选项A内容",
|
||||||
"B": "选项B内容",
|
"选项B内容",
|
||||||
"C": "选项C内容",
|
"选项C内容",
|
||||||
"D": "选项D内容"
|
"选项D内容"
|
||||||
|
],
|
||||||
|
"answer": ["A"]
|
||||||
},
|
},
|
||||||
"answer": "A"
|
{
|
||||||
|
"id": 2,
|
||||||
|
"type": "multiple",
|
||||||
|
"question": "题目内容",
|
||||||
|
"options": [
|
||||||
|
"选项A内容",
|
||||||
|
"选项B内容",
|
||||||
|
"选项C内容",
|
||||||
|
"选项D内容"
|
||||||
|
],
|
||||||
|
"answer": ["A", "C"]
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
|
||||||
"""
|
"""
|
||||||
27
database.py
Normal file
27
database.py
Normal 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
0
id_generator/__init__.py
Normal file
38
id_generator/generator.py
Normal file
38
id_generator/generator.py
Normal 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
134
id_generator/idregister.py
Normal 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
43
id_generator/options.py
Normal 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
20
id_generator/snowflake.py
Normal 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
|
||||||
147
id_generator/snowflake_m1.py
Normal file
147
id_generator/snowflake_m1.py
Normal 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
28
main.py
@@ -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.middleware.cors import CORSMiddleware
|
||||||
from starlette.responses import StreamingResponse
|
from starlette.responses import StreamingResponse
|
||||||
|
|
||||||
from models.agent import QueryRequest
|
from database import get_db
|
||||||
from service import query_agent
|
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 storage import upload_rustfs
|
||||||
from ocr import ocr_from_url
|
from ocr import ocr_from_url
|
||||||
|
|
||||||
@@ -34,3 +40,19 @@ def generate(query: QueryRequest):
|
|||||||
query_agent(query),
|
query_agent(query),
|
||||||
media_type="text/event-stream"
|
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
43
models/base.py
Normal 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
33
models/library.py
Normal 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"
|
||||||
|
)
|
||||||
@@ -8,3 +8,4 @@ langchain-core~=1.5.1
|
|||||||
langchain-openai~=1.4.1
|
langchain-openai~=1.4.1
|
||||||
starlette~=1.3.1
|
starlette~=1.3.1
|
||||||
pydantic~=2.13.4
|
pydantic~=2.13.4
|
||||||
|
SQLAlchemy~=2.0.51
|
||||||
37
schemas/library.py
Normal file
37
schemas/library.py
Normal 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
|
||||||
@@ -1,7 +1,7 @@
|
|||||||
from langchain_core.messages import AIMessageChunk, HumanMessage
|
from langchain_core.messages import AIMessageChunk, HumanMessage
|
||||||
|
|
||||||
from agent.study import study_agent
|
from agent.study import study_agent
|
||||||
from models.agent import QueryRequest
|
from schemas.agent import QueryRequest
|
||||||
|
|
||||||
|
|
||||||
def query_agent(query: QueryRequest):
|
def query_agent(query: QueryRequest):
|
||||||
85
service/library_service.py
Normal file
85
service/library_service.py
Normal 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
13
utils/common.py
Normal 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)
|
||||||
Reference in New Issue
Block a user