This commit is contained in:
dni ⚡
2024-10-10 09:07:35 +02:00
parent cae81b61a3
commit d813544bf4
+36 -4
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import asyncio import asyncio
import datetime import datetime
import json
import os import os
import re import re
import time import time
@@ -144,28 +145,40 @@ class Connection(Compat):
clean_values[key] = raw_value clean_values[key] = raw_value
return clean_values return clean_values
async def fetchall(self, query: str, values: Optional[dict] = None) -> list[dict]: async def fetchall(
self, query: str, values: Optional[dict] = None, model: Optional[TModel] = None
) -> list[TModel]:
params = self.rewrite_values(values) if values else {} params = self.rewrite_values(values) if values else {}
result = await self.conn.execute(text(self.rewrite_query(query)), params) result = await self.conn.execute(text(self.rewrite_query(query)), params)
row = result.mappings().all() row = result.mappings().all()
result.close() result.close()
if not row:
return []
if model:
return [_dict_to_model(r, model) for r in row]
return row return row
async def fetchone(self, query: str, values: Optional[dict] = None) -> dict: async def fetchone(
self, query: str, values: Optional[dict] = None, model: Optional[TModel] = None
) -> TModel:
params = self.rewrite_values(values) if values else {} params = self.rewrite_values(values) if values else {}
result = await self.conn.execute(text(self.rewrite_query(query)), params) result = await self.conn.execute(text(self.rewrite_query(query)), params)
row = result.mappings().first() row = result.mappings().first()
result.close() result.close()
if model and row:
return _dict_to_model(row, model)
return row return row
async def update(self, table_name: str, model: BaseModel, where: str = "id = :id"): async def update(self, table_name: str, model: BaseModel, where: str = "id = :id"):
await self.conn.execute( await self.conn.execute(
text(update_query(table_name, model, where)), model.dict() text(update_query(table_name, model, where)), _model_to_dict(model)
) )
await self.conn.commit() await self.conn.commit()
async def insert(self, table_name: str, model: BaseModel): async def insert(self, table_name: str, model: BaseModel):
await self.conn.execute(text(insert_query(table_name, model)), model.dict()) await self.conn.execute(
text(insert_query(table_name, model)), _model_to_dict(model)
)
await self.conn.commit() await self.conn.commit()
async def fetch_page( async def fetch_page(
@@ -569,3 +582,22 @@ def update_query(
fields.append(f"{field} = {placeholder}") fields.append(f"{field} = {placeholder}")
query = ", ".join(fields) query = ", ".join(fields)
return f"UPDATE {table_name} SET {query} {where}" return f"UPDATE {table_name} SET {query} {where}"
def _model_to_dict(model: BaseModel) -> dict:
_dict = model.dict()
for key, value in _dict.items():
if key.startswith("_"):
continue
type_ = model.__fields__[key].type_
if type_ == BaseModel:
_dict[key] = json.dumps(value.dict())
return _dict
def _dict_to_model(_dict: dict, model: TModel) -> TModel:
for key, value in _dict.items():
type_ = model.__fields__[key].type_
if type_ is BaseModel:
_dict[key] = json.loads(value)
return model.construct(**_dict)