完善 aiosqlite 异步数据库驱动

This commit is contained in:
lan-air
2022-12-11 17:41:54 +08:00
parent d25a92199f
commit e5d18a5bf3
3 changed files with 15 additions and 18 deletions
-1
View File
@@ -5,7 +5,6 @@ from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.ext.asyncio import create_async_engine from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.ext.asyncio.session import AsyncSession from sqlalchemy.ext.asyncio.session import AsyncSession
engine = create_async_engine("sqlite+aiosqlite:///database.db") engine = create_async_engine("sqlite+aiosqlite:///database.db")
Base = declarative_base() Base = declarative_base()
+10 -16
View File
@@ -3,22 +3,15 @@ import os
import uuid import uuid
import threading import threading
import random import random
from fastapi import FastAPI, Depends, UploadFile, Form, File from fastapi import FastAPI, Depends, UploadFile, Form, File
from starlette.requests import Request from starlette.requests import Request
from starlette.responses import HTMLResponse, FileResponse from starlette.responses import HTMLResponse, FileResponse
import random
from starlette.staticfiles import StaticFiles from starlette.staticfiles import StaticFiles
from sqlalchemy import or_, select, update, delete, create_engine from sqlalchemy import or_, select, update, delete
from sqlalchemy import select, update, delete
from sqlalchemy.ext.asyncio.session import AsyncSession from sqlalchemy.ext.asyncio.session import AsyncSession
from database import engine, get_session, Base, Codes from database import get_session, Codes
engine = create_engine('sqlite:///database.db', connect_args={"check_same_thread": False})
Base.metadata.create_all(bind=engine)
app = FastAPI() app = FastAPI()
if not os.path.exists('./static'): if not os.path.exists('./static'):
@@ -137,13 +130,14 @@ def ip_error(ip):
@app.get('/select') @app.get('/select')
async def get_file(code: str, db: Session = Depends(get_db)): async def get_file(code: str, s: AsyncSession = Depends(get_session)):
file = db.query(database.Codes).filter(database.Codes.code == code).first() query = select(Codes).where(Codes.code == code)
if file: info = (await s.execute(query)).scalars().first()
if file.type == 'text': if info:
return {'code': code, 'msg': '查询成功', 'data': file.text} if info.type == 'text':
return {'code': code, 'msg': '查询成功', 'data': info.text}
else: else:
return FileResponse('.' + file.text, filename=file.name) return FileResponse('.' + info.text, filename=info.name)
else: else:
return {'code': 404, 'msg': '口令不存在'} return {'code': 404, 'msg': '口令不存在'}
@@ -182,7 +176,7 @@ async def share(text: str = Form(default=None), style: str = Form(default='2'),
query = select(Codes).where(or_(Codes.exp_time < datetime.datetime.now(), Codes.count == 0)) query = select(Codes).where(or_(Codes.exp_time < datetime.datetime.now(), Codes.count == 0))
exps = (await s.execute(query)).scalars().all() exps = (await s.execute(query)).scalars().all()
threading.Thread(target=delete_file, args=([[{'type': old.type, 'text': old.text}] for old in exps],)).start() 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] exps_ids = [exp.id for exp in exps]
query = delete(Codes).where(Codes.id.in_(exps_ids)) query = delete(Codes).where(Codes.id.in_(exps_ids))
await s.execute(query) await s.execute(query)
+5 -1
View File
@@ -1,3 +1,7 @@
fastapi[all]==0.88.0 fastapi==0.88.0
aiosqlite==0.17.0 aiosqlite==0.17.0
SQLAlchemy==1.4.44 SQLAlchemy==1.4.44
python-multipart==0.0.5
uvicorn==0.15.0
greenlet==2.0.1
starlette~=0.22.0