Skip to content

Commit ee4d5e1

Browse files
committed
clean up and make it go brrrrr
1 parent 14ab52a commit ee4d5e1

2 files changed

Lines changed: 24 additions & 43 deletions

File tree

litequery/core.py

Lines changed: 24 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22
import os
33
import re
44
import sqlite3
5-
from collections import OrderedDict
65
from collections.abc import Iterator
76
from contextlib import contextmanager
87
from dataclasses import dataclass
@@ -12,10 +11,6 @@
1211

1312
from litequery.config import Config, get_config
1413

15-
_iso8601_pattern = re.compile(
16-
r"^\d{4}-\d{2}-\d{2}[T ]\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?$"
17-
)
18-
1914

2015
class Op(str, Enum):
2116
SELECT = ""
@@ -35,23 +30,28 @@ class Query:
3530

3631
class Row:
3732
def __init__(self, columns: list[str], values: list[Any]):
38-
self._data = OrderedDict()
39-
self._values: list[Any] = []
40-
41-
for col, val in zip(columns, values):
42-
if isinstance(val, str) and _iso8601_pattern.match(val):
43-
try:
44-
val = datetime.fromisoformat(val)
45-
except ValueError:
46-
pass
47-
self._values.append(val)
48-
self._data[col] = val
33+
if len(set(columns)) != len(columns):
34+
dups = [c for c in columns if columns.count() > 1]
35+
raise ValueError(f"Duplicate columns: {set(dups)}. Use AS to alias.")
36+
37+
self._values = tuple(
38+
self._parse_datetime(v) if isinstance(v, str) else v for v in values
39+
)
40+
self._index = {c: i for i, c in enumerate(columns)}
41+
42+
def _parse_datetime(self, value: str):
43+
if 19 <= len(value) <= 32 and value[0].isdigit():
44+
try:
45+
return datetime.fromisoformat(value)
46+
except ValueError:
47+
pass
48+
return value
4949

5050
def _available_columns(self) -> str:
51-
return ", ".join([f"'{c}'" for c in self._data.keys()])
51+
return ", ".join([f"'{c}'" for c in self._index.keys()])
5252

5353
def __repr__(self) -> str:
54-
items = [f"{col}={repr(val)}" for col, val in self._data.items()]
54+
items = [f"{col}={self._values[idx]!r}" for col, idx in self._index.items()]
5555
return f"{self.__class__.__name__}({', '.join(items)})"
5656

5757
def __getitem__(self, key: int | str) -> Any:
@@ -64,7 +64,7 @@ def __getitem__(self, key: int | str) -> Any:
6464
f"can't access index {key}"
6565
)
6666
try:
67-
return self._data[key]
67+
return self._values[self._index[key]]
6868
except KeyError:
6969
raise KeyError(
7070
f"No column '{key}' found. Available: {self._available_columns()}"
@@ -78,16 +78,10 @@ def __getattr__(self, name: str) -> Any:
7878
raise error
7979

8080
try:
81-
return self._data[name]
81+
return self._values[self._index[name]]
8282
except KeyError:
8383
raise error
8484

85-
def __setitem__(self, key: int | str, value: Any) -> None:
86-
raise TypeError("Row assignment not supported")
87-
88-
def __contains__(self, name: str) -> bool:
89-
return name in self._data
90-
9185
def __len__(self) -> int:
9286
return len(self._values)
9387

@@ -97,27 +91,18 @@ def __iter__(self) -> Iterator[Any]:
9791
def __eq__(self, other) -> bool:
9892
if not isinstance(other, Row):
9993
return False
100-
return self._data == other._data
101-
102-
def keys(self) -> list[str]:
103-
return list(self._data.keys())
104-
105-
def values(self) -> list[Any]:
106-
return list(self._data.values())
107-
108-
def items(self) -> list[tuple[str, Any]]:
109-
return list(self._data.items())
94+
return self._values == other._values
11095

11196
def to_dict(self) -> dict:
112-
return dict(self._data)
97+
return dict(zip(self._index.keys(), self._values))
11398

11499
def into(self, cls):
115100
return cls(**self.to_dict())
116101

117102

118103
class Rows(list):
119104
def into(self, cls):
120-
return [cls(**row.to_dict()) for row in self]
105+
return Rows([cls(**row.to_dict()) for row in self])
121106

122107

123108
def parse_file_queries(file_path):
@@ -127,7 +112,7 @@ def parse_file_queries(file_path):
127112

128113
queries = []
129114
op_pattern = "|".join("\\" + "\\".join(list(op.value)) for op in Op if op.value)
130-
pattern = rf"^([a-z_][a-z0-9_-]*)({op_pattern})?$"
115+
pattern = rf"^([a-z_][a-z0-9_]*)({op_pattern})?$"
131116
for query_name, sql in raw_queries:
132117
match = re.match(pattern, query_name)
133118
if not match:
@@ -188,9 +173,6 @@ def _create_methods(self, queries: list[Query]):
188173
for query in queries:
189174
setattr(self, query.name, self._create_method(query))
190175

191-
def _create_method(self, query):
192-
raise NotImplementedError("This method should be overridden!")
193-
194176
def _execute_query(self, conn: sqlite3.Connection, query: Query, kwargs: dict):
195177
cursor = conn.cursor()
196178
cursor.execute(query.sql, kwargs)

tests/test_core.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -91,7 +91,6 @@ class User:
9191
assert isinstance(user, User)
9292

9393
users = lq.get_all_users().into(User)
94-
breakpoint()
9594
assert isinstance(users[0], User)
9695

9796

0 commit comments

Comments
 (0)