munch-ease-backend/products/db.py

97 lines
3 KiB
Python

from typing import AsyncIterator, List, ClassVar
from pydantic import BaseModel
import json
class Product(BaseModel):
KEYS: ClassVar[List[str]] = ['id', 'product_id', 'link', 'name', 'quantity', 'unit', 'img_small', 'img_large']
NON_INSERT_KEYS: ClassVar[List[str]] = ['id']
id: int = -1
product_id: str
link: str
name: str
quantity: int
unit: str
img_small: str
img_large: str
async def create(conn):
await conn.execute('''
CREATE TABLE IF NOT EXISTS Product (
id INTEGER PRIMARY KEY,
product_id TEXT UNIQUE,
link TEXT,
name TEXT,
quantity INTEGER,
unit TEXT,
img_small TEXT,
img_large TEXT,
raw_data TEXT
);''')
await conn.execute('''
CREATE TABLE IF NOT EXISTS ProductTag (
food_item_id INTEGER,
tag TEXT COLLATE NOCASE,
PRIMARY KEY (food_item_id, tag),
FOREIGN KEY (food_item_id) REFERENCES Product(id)
);''')
async def find_product_by_tag(conn, tag: str) -> AsyncIterator[Product]:
async with conn.execute(f'''
SELECT {','.join(Product.KEYS)} FROM Product
WHERE id IN (
SELECT food_item_id FROM ProductTag
WHERE tag = ?
)
''', (tag,)) as cursor:
async for row in cursor:
yield Product(**{k:v for k,v in zip(Product.KEYS, row)})
async def find_product_by_id(conn, product_id: str) -> Product:
async with conn.execute(f'''
SELECT {','.join(Product.KEYS)} FROM Product
WHERE id = ?
LIMIT 1
''', (product_id,)) as cursor:
async for row in cursor:
return Product(**{k:v for k,v in zip(Product.KEYS, row)})
async def find_product_by_product_id(conn, product_id: str) -> Product:
async with conn.execute(f'''
SELECT {','.join(Product.KEYS)} FROM Product
WHERE product_id = ?
LIMIT 1
''', (product_id,)) as cursor:
async for row in cursor:
return Product(**{k:v for k,v in zip(Product.KEYS, row)})
async def insert_product(conn, product: Product, data: dict):
insert_keys = [k for k in Product.KEYS if k not in Product.NON_INSERT_KEYS]
insert_values = [getattr(product, k) for k in insert_keys]
async with conn.execute(f'''
INSERT INTO Product ({','.join(insert_keys)}, raw_data)
VALUES ({','.join(['?'] * len(insert_keys))}, ?)
''', (*insert_values, json.dumps(data))) as cursor:
product.id = cursor.lastrowid
await conn.commit()
async def add_tag(conn, product: Product, tag: str):
await conn.execute('''
INSERT INTO ProductTag (food_item_id, tag)
VALUES (?, ?)
''', (product.id, tag))
await conn.commit()
async def get_tags(conn, product: Product) -> AsyncIterator[str]:
async with conn.execute('''
SELECT tag FROM ProductTag
WHERE food_item_id = ?
''', (product.id,)) as cursor:
async for row in cursor:
yield row[0]