Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 6 additions & 6 deletions ddss/egg.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = []
Expand Down
8 changes: 4 additions & 4 deletions ddss/egraph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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] = {}

Expand Down
2 changes: 1 addition & 1 deletion ddss/input.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down
8 changes: 5 additions & 3 deletions ddss/main.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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,
Expand All @@ -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
Expand Down
20 changes: 11 additions & 9 deletions ddss/orm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down
12 changes: 6 additions & 6 deletions ddss/output.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading