retronews

a featureful fork of the luke8086/retronews hn+lobste.rs tui
Log | Files | Refs | README | LICENSE

commit 61b3fe656b0c2cd86f093d5cd9f0ed71480a7f64
parent 56a46f5661d9c71aff3dbcc4f7434e376ea4c86c
Author: luke8086 <55237178+luke8086@users.noreply.github.com>
Date:   Sun, 31 Jul 2022 12:54:40 +0000

Use type aliases for curses windows and sqlite3 connections

Diffstat:
Mretronews.py | 29++++++++++++++++-------------
1 file changed, 16 insertions(+), 13 deletions(-)

diff --git a/retronews.py b/retronews.py @@ -24,6 +24,8 @@ from functools import partial from textwrap import wrap from typing import Any, Callable, Generator, Optional, TypedDict, TypeVar, Union +from typing_extensions import TypeAlias + KEY_BINDINGS = { ord("q"): lambda app: cmd_quit(app), ord("?"): lambda app: cmd_help(app), @@ -76,10 +78,11 @@ Press any key to continue... """ REQUEST_TIMEOUT = 10 +QUOTE_REX = re.compile(r"^(> ?)+") T = TypeVar("T") - -QUOTE_REX = re.compile(r"^(> ?)+") +Window: TypeAlias = "curses._CursesWindow" +DB: TypeAlias = "sqlite3.Connection" class Colors: @@ -182,9 +185,9 @@ class Layout: @dataclasses.dataclass class AppState: - screen: "curses._CursesWindow" + screen: Window colors: Colors - db: sqlite3.Connection + db: DB group: Group messages: list[Message] = dataclasses.field(default_factory=list) messages_by_id: dict[str, Message] = dataclasses.field(default_factory=dict) @@ -463,7 +466,7 @@ def cmd_unknown(app: AppState) -> None: app.flash = "Unknown key" -def db_init(path: str) -> sqlite3.Connection: +def db_init(path: str) -> DB: path = os.path.expanduser(path) create_table_sql = """ CREATE TABLE IF NOT EXISTS messages ( @@ -482,14 +485,14 @@ def db_init(path: str) -> sqlite3.Connection: return db -def db_save_message(db: sqlite3.Connection, message: Message) -> None: +def db_save_message(db: DB, message: Message) -> None: sql = """INSERT OR REPLACE INTO messages (id, thread_id, flags) VALUES (?, ?, ?)""" flags_json = json.dumps(dataclasses.asdict(message.flags)) db.execute(sql, (message.msg_id, message.thread_id, flags_json)) db.commit() -def db_load_message_flags(db: sqlite3.Connection, messages_by_id: dict[str, Message]) -> None: +def db_load_message_flags(db: DB, messages_by_id: dict[str, Message]) -> None: message_ids = list(messages_by_id.keys()) sql = f"SELECT * FROM messages WHERE id IN ({','.join('?' for _ in message_ids)})" @@ -498,7 +501,7 @@ def db_load_message_flags(db: sqlite3.Connection, messages_by_id: dict[str, Mess messages_by_id[row["id"]].flags = MessageFlags(**flags) -def db_load_read_comments(db: sqlite3.Connection, messages_by_id: dict[str, Message]) -> None: +def db_load_read_comments(db: DB, messages_by_id: dict[str, Message]) -> None: threads_by_id = {msg.msg_id: msg for msg in messages_by_id.values() if msg.is_thread} thread_ids = list(threads_by_id.keys()) @@ -513,7 +516,7 @@ def db_load_read_comments(db: sqlite3.Connection, messages_by_id: dict[str, Mess threads_by_id[row["thread_id"]].read_comments = row["count"] -def db_load_starred_thread_ids(db: sqlite3.Connection, page: int = 1) -> list[str]: +def db_load_starred_thread_ids(db: DB, page: int = 1) -> list[str]: page_size = 30 offset = (page - 1) * page_size sql = """ @@ -827,7 +830,7 @@ def app_render(app: AppState) -> None: app.screen.refresh() -def app_init(screen: "curses._CursesWindow", db: sqlite3.Connection) -> AppState: +def app_init(screen: Window, db: DB) -> AppState: curses.curs_set(0) curses.use_default_colors() @@ -945,7 +948,7 @@ def group_advance_page(group: Group, offset: int = 1) -> Group: return group_set_page(group, page=max(1, group.page + offset)) -def group_search_starred_threads(db: sqlite3.Connection, page: int = 1) -> list[Message]: +def group_search_starred_threads(db: DB, page: int = 1) -> list[Message]: thread_ids = db_load_starred_thread_ids(db, page) threads_by_provider: dict[str, list[str]] = {} threads = [] @@ -962,7 +965,7 @@ def group_search_starred_threads(db: sqlite3.Connection, page: int = 1) -> list[ return threads -def group_search_threads(group: Group, db: sqlite3.Connection) -> list[Message]: +def group_search_threads(group: Group, db: DB) -> list[Message]: if group.provider == "hn": return hn_search_threads(group.name, group.page) elif group.provider == "hn-new": @@ -992,7 +995,7 @@ def logging_init(path: Optional[str]) -> None: logging.debug("Session started") -def main(screen: "curses._CursesWindow", db: sqlite3.Connection) -> None: +def main(screen: Window, db: DB) -> None: app = app_init(screen, db) while True: