fixup!
This commit is contained in:
+36
-4
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user