数据库操作改为异步

This commit is contained in:
veoco
2022-12-11 14:04:25 +08:00
parent 195c6f9ae4
commit f4e6dfef0d
2 changed files with 57 additions and 49 deletions
+12 -6
View File
@@ -1,17 +1,23 @@
import datetime
from sqlalchemy import create_engine, DateTime
from sqlalchemy.orm import sessionmaker
from sqlalchemy import Boolean, Column, Integer, String, DateTime
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy import Boolean, Column, Integer, String
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.ext.asyncio.session import AsyncSession
engine = create_async_engine("sqlite+aiosqlite:///database.db")
engine = create_engine('sqlite:///database.db', connect_args={"check_same_thread": False})
Base = declarative_base()
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
async def get_session():
async with AsyncSession(engine, expire_on_commit=False) as s:
yield s
class Codes(Base):
__tablename__ = 'codes'
__tablename__ = "codes"
id = Column(Integer, primary_key=True, index=True)
code = Column(String(10), unique=True, index=True)
key = Column(String(30), unique=True, index=True)
+45 -43
View File
@@ -2,19 +2,23 @@ import datetime
import os
import uuid
import threading
from fastapi import FastAPI, Depends, UploadFile, Form, File
from sqlalchemy import or_
from sqlalchemy.orm import Session
from starlette.requests import Request
from starlette.responses import HTMLResponse
import random
from fastapi import FastAPI, Depends, UploadFile, Form, File
from starlette.requests import Request
from starlette.responses import HTMLResponse
from starlette.staticfiles import StaticFiles
import database
from database import engine, SessionLocal, Base
from sqlalchemy import or_, select, update, delete, create_engine
from sqlalchemy import select, update, delete
from sqlalchemy.ext.asyncio.session import AsyncSession
from database import engine, get_session, Base, Codes
engine = create_engine('sqlite:///database.db', connect_args={"check_same_thread": False})
Base.metadata.create_all(bind=engine)
app = FastAPI()
if not os.path.exists('./static'):
os.makedirs('./static')
@@ -58,17 +62,9 @@ def delete_file(files):
os.remove('.' + file['text'])
def get_db():
db = SessionLocal()
try:
yield db
finally:
db.close()
def get_code(db: Session = Depends(get_db)):
async def get_code(s: AsyncSession):
code = random.randint(10000, 99999)
while db.query(database.Codes).filter(database.Codes.code == code).first():
while (await s.execute(select(Codes.id).where(Codes.code == code))).scalar():
code = random.randint(10000, 99999)
return str(code)
@@ -94,21 +90,23 @@ async def admin():
@app.post(f'/{admin_address}')
async def admin_post(request: Request, db: Session = Depends(get_db)):
async def admin_post(request: Request, s: AsyncSession = Depends(get_session)):
if request.headers.get('pwd') == admin_password:
codes = db.query(database.Codes).all()
query = select(Codes)
codes = (await s.execute(query)).scalars().all()
return {'code': 200, 'msg': '查询成功', 'data': codes}
else:
return {'code': 404, 'msg': '密码错误'}
@app.delete(f'/{admin_address}')
async def admin_delete(request: Request, code: str, db: Session = Depends(get_db)):
async def admin_delete(request: Request, code: str, s: AsyncSession = Depends(get_session)):
if request.headers.get('pwd') == admin_password:
file = db.query(database.Codes).filter(database.Codes.code == code).first()
query = select(Codes).where(Codes.code == code)
file = (await s.execute(query)).scalars().first()
threading.Thread(target=delete_file, args=([{'type': file.type, 'text': file.text}],)).start()
db.delete(file)
db.commit()
await s.delete(file)
await s.commit()
return {'code': 200, 'msg': '删除成功'}
else:
return {'code': 404, 'msg': '密码错误'}
@@ -138,20 +136,24 @@ def ip_error(ip):
@app.post('/')
async def index(request: Request, code: str, db: Session = Depends(get_db)):
async def index(request: Request, code: str, s: AsyncSession = Depends(get_session)):
ip = request.client.host
if not check_ip(ip):
return {'code': 404, 'msg': '错误次数过多,请稍后再试'}
info = db.query(database.Codes).filter(database.Codes.code == code).first()
query = select(Codes).where(Codes.code == code)
info = (await s.execute(query)).scalars().first()
if not info:
return {'code': 404, 'msg': f'取件码错误,错误{error_count - ip_error(ip)}次将被禁止10分钟'}
if info.exp_time < datetime.datetime.now() or info.count == 0:
threading.Thread(target=delete_file, args=([{'type': info.type, 'text': info.text}],)).start()
db.delete(info)
db.commit()
await s.delete(info)
await s.commit()
return {'code': 404, 'msg': '取件码已过期,请联系寄件人'}
info.count -= 1
db.commit()
count = info.count - 1
query = update(Codes).where(Codes.id == info.id).values(count=count)
await s.execute(query)
await s.commit()
return {
'code': 200,
'msg': '取件成功,请点击""查看',
@@ -161,17 +163,17 @@ async def index(request: Request, code: str, db: Session = Depends(get_db)):
@app.post('/share')
async def share(text: str = Form(default=None), style: str = Form(default='2'), value: int = Form(default=1),
file: UploadFile = File(default=None), db: Session = Depends(get_db)):
exps = db.query(database.Codes).filter(
or_(
database.Codes.exp_time < datetime.datetime.now(),
database.Codes.count == 0
)
)
threading.Thread(target=delete_file, args=([[{'type': old.type, 'text': old.text}] for old in exps.all()],)).start()
exps.delete()
db.commit()
code = get_code(db)
file: UploadFile = File(default=None), s: AsyncSession = Depends(get_session)):
query = select(Codes).where(or_(Codes.exp_time < datetime.datetime.now(), Codes.count == 0))
exps = (await s.execute(query)).scalars().all()
threading.Thread(target=delete_file, args=([[{'type': old.type, 'text': old.text}] for old in exps],)).start()
exps_ids = [exp.id for exp in exps]
query = delete(Codes).where(Codes.id.in_(exps_ids))
await s.execute(query)
await s.commit()
code = await get_code(s)
if style == '2':
if value > 7:
return {'code': 404, 'msg': '最大有效天数为7天'}
@@ -192,7 +194,7 @@ async def share(text: str = Form(default=None), style: str = Form(default='2'),
return {'code': 404, 'msg': '文件过大'}
else:
size, _text, _type, name = len(text), text, 'text', '文本分享'
info = database.Codes(
info = Codes(
code=code,
text=_text,
size=size,
@@ -202,8 +204,8 @@ async def share(text: str = Form(default=None), style: str = Form(default='2'),
exp_time=exp_time,
key=key
)
db.add(info)
db.commit()
s.add(info)
await s.commit()
return {
'code': 200,
'msg': '分享成功,请点击文件箱查看取件码',