diff --git a/ddss/egg.py b/ddss/egg.py index fa1f6ca..9bf69db 100644 --- a/ddss/egg.py +++ b/ddss/egg.py @@ -15,12 +15,12 @@ async def main(session: async_sessionmaker[AsyncSession]) -> None: while True: async with session() as sess: - for i in await sess.scalars(select(Ideas).where(Ideas.id > max_idea)): - max_idea = max(max_idea, i.id) - pool.append(Rule(i.data)) - for i in await sess.scalars(select(Facts).where(Facts.id > max_fact)): - max_fact = max(max_fact, i.id) - search.add(Rule(i.data)) + for idea in await sess.scalars(select(Ideas).where(Ideas.id > max_idea)): + max_idea = max(max_idea, idea.id) + pool.append(Rule(idea.data)) + for fact in await sess.scalars(select(Facts).where(Facts.id > max_fact)): + max_fact = max(max_fact, fact.id) + search.add(Rule(fact.data)) search.rebuild() tasks = [] next_pool = [] diff --git a/ddss/egraph.py b/ddss/egraph.py index b7a7871..946d58a 100644 --- a/ddss/egraph.py +++ b/ddss/egraph.py @@ -11,12 +11,12 @@ def _build_term_to_rule(data: Term) -> Rule: def _extract_lhs_rhs_from_rule(data: Rule) -> tuple[Term, Term] | None: if len(data) != 0: - return + return None term = data.conclusion.term if not isinstance(term, List): - return + return None if not (len(term) == 4 or str(term[0]) == "binary" or str(term[1]) == "=="): - return + return None lhs = term[2] rhs = term[3] return lhs, rhs @@ -27,7 +27,7 @@ def _build_lhs_rhs_to_term(lhs: Term, rhs: Term) -> Term: class _EGraph: - def __init__(self): + def __init__(self) -> None: self.core = EGraph() self.mapping: dict[Term, EClassId] = {} diff --git a/ddss/input.py b/ddss/input.py index 543afb0..f076e26 100644 --- a/ddss/input.py +++ b/ddss/input.py @@ -9,7 +9,7 @@ async def main(session: async_sessionmaker[AsyncSession]) -> None: try: - prompt = PromptSession() + prompt: PromptSession[str] = PromptSession() while True: try: with patch_stdout(): diff --git a/ddss/main.py b/ddss/main.py index 798c1f5..cf9e882 100644 --- a/ddss/main.py +++ b/ddss/main.py @@ -1,7 +1,7 @@ import asyncio import tempfile import pathlib -from typing import Annotated, Optional, Awaitable +from typing import Annotated, Optional, Callable, Coroutine import tyro from sqlalchemy.ext.asyncio import async_sessionmaker, AsyncSession from .orm import initialize_database @@ -13,7 +13,7 @@ from .dump import main as dump from .chain import main as chain -component_map: dict[str, callable[[async_sessionmaker[AsyncSession]], Awaitable[None]]] = { +component_map: dict[str, Callable[[async_sessionmaker[AsyncSession]], Coroutine[None, None, None]]] = { "search": search, "egg": egg, "input": input, @@ -29,7 +29,9 @@ async def run(addr: str, components: list[str]) -> None: try: try: - coroutines = [component_map[component](session) for component in components] + coroutines: list[Coroutine[None, None, None]] = [ + component_map[component](session) for component in components + ] except KeyError as e: print(f"error: unsupported component: {str(e)}") return diff --git a/ddss/orm.py b/ddss/orm.py index 6608380..caeb9f0 100644 --- a/ddss/orm.py +++ b/ddss/orm.py @@ -5,9 +5,9 @@ from sqlalchemy.exc import IntegrityError from sqlalchemy import Integer, Text from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column -from sqlalchemy.dialects.sqlite import insert as sqlite_insert -from sqlalchemy.dialects.mysql import insert as mysql_insert -from sqlalchemy.dialects.postgresql import insert as postgresql_insert +from sqlalchemy.dialects.sqlite import insert as sqlite_insert, Insert as sqlite_insert_t +from sqlalchemy.dialects.mysql import insert as mysql_insert, Insert as mysql_insert_t +from sqlalchemy.dialects.postgresql import insert as postgresql_insert, Insert as postgresql_insert_t class Base(DeclarativeBase): @@ -37,14 +37,16 @@ async def initialize_database(addr: str) -> tuple[AsyncEngine, async_sessionmake async def insert_or_ignore(sess: AsyncSession, model: type[Base], data: str, locks=defaultdict(asyncio.Lock)) -> None: match sess.bind.dialect.name: case "sqlite": - statement = sqlite_insert(model).values(data=data).on_conflict_do_nothing() - await asyncio.shield(sess.execute(statement)) + sqlite_statement: sqlite_insert_t = sqlite_insert(model).values(data=data).on_conflict_do_nothing() + await asyncio.shield(sess.execute(sqlite_statement)) case "mysql" | "mariadb": - statement = mysql_insert(model).values(data=data).prefix_with("IGNORE") - await asyncio.shield(sess.execute(statement)) + mysql_statement: mysql_insert_t = mysql_insert(model).values(data=data).prefix_with("IGNORE") + await asyncio.shield(sess.execute(mysql_statement)) case "postgresql": - statement = postgresql_insert(model).values(data=data).on_conflict_do_nothing() - await asyncio.shield(sess.execute(statement)) + postgresql_statement: postgresql_insert_t = ( + postgresql_insert(model).values(data=data).on_conflict_do_nothing() + ) + await asyncio.shield(sess.execute(postgresql_statement)) case _: async with locks[id(sess.bind)]: try: diff --git a/ddss/output.py b/ddss/output.py index 903642d..5a03b66 100644 --- a/ddss/output.py +++ b/ddss/output.py @@ -12,12 +12,12 @@ async def main(session: async_sessionmaker[AsyncSession]) -> None: while True: async with session() as sess: - for i in await sess.scalars(select(Ideas).where(Ideas.id > max_idea)): - max_idea = max(max_idea, i.id) - print("idea:", unparse(i.data)) - for i in await sess.scalars(select(Facts).where(Facts.id > max_fact)): - max_fact = max(max_fact, i.id) - print("fact:", unparse(i.data)) + for idea in await sess.scalars(select(Ideas).where(Ideas.id > max_idea)): + max_idea = max(max_idea, idea.id) + print("idea:", unparse(idea.data)) + for fact in await sess.scalars(select(Facts).where(Facts.id > max_fact)): + max_fact = max(max_fact, fact.id) + print("fact:", unparse(fact.data)) await sess.commit() await asyncio.sleep(0)