Skip to content

Commit 43e0b3a

Browse files
committed
feat: enhance batch processing with strict result count checks and null handling
1 parent 8a066e1 commit 43e0b3a

10 files changed

Lines changed: 315 additions & 47 deletions

docs/reference.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -171,7 +171,7 @@ Both sync and async connections support:
171171

172172
## Project Structure
173173

174-
```
174+
```text
175175
qql/
176176
├── pyproject.toml # Package config; installs the `qql` CLI command
177177
├── src/

src/qql/async_connection.py

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -190,22 +190,36 @@ async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseExc
190190
if not self._queries:
191191
return
192192
results = await self.connection.run_queries_batch(self._queries)
193-
for proxy, res in zip(self._proxies, results):
193+
for proxy, res in zip(self._proxies, results, strict=False):
194194
proxy._resolve(res)
195+
if len(results) != len(self._proxies):
196+
error = RuntimeError(
197+
"Batch result count mismatch: "
198+
f"expected {len(self._proxies)}, got {len(results)}"
199+
)
200+
for proxy in self._proxies[len(results):]:
201+
proxy._reject(error)
202+
raise error
195203

196204

197205
class AsyncOperationProxy:
198206
"""Proxy handle that resolves to an ExecutionResult after QQLAsyncBatch exits."""
199207

200208
def __init__(self) -> None:
201209
self._result: ExecutionResult | None = None
210+
self._exception: RuntimeError | None = None
202211

203212
def _resolve(self, result: ExecutionResult) -> None:
204213
self._result = result
205214

215+
def _reject(self, exception: RuntimeError) -> None:
216+
self._exception = exception
217+
206218
@property
207219
def result(self) -> ExecutionResult:
208220
"""The resolved ExecutionResult."""
221+
if self._exception is not None:
222+
raise self._exception
209223
if self._result is None:
210224
raise RuntimeError("AsyncBatch has not been executed yet.")
211225
return self._result

src/qql/async_executor.py

Lines changed: 30 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -138,12 +138,16 @@ async def _ensure_collection(
138138
vector_size: int,
139139
topology: CollectionTopology,
140140
explicit_vector: str | None,
141-
) -> None:
141+
) -> CollectionTopology:
142142
if topology.exists:
143143
info = await self._client.get_collection(name)
144144
vectors = info.config.params.vectors # type: ignore[union-attr]
145+
sparse_vectors = info.config.params.sparse_vectors or {}
146+
current_topology = CollectionTopology(
147+
**collection_topology_kwargs(vectors, sparse_vectors)
148+
)
145149
if isinstance(vectors, dict):
146-
vector_name = topology.dense_using(explicit_vector)
150+
vector_name = current_topology.dense_using(explicit_vector)
147151
if vector_name is None:
148152
raise QQLRuntimeError("Collection has no dense vector")
149153
vector_config = vectors[vector_name]
@@ -164,12 +168,12 @@ async def _ensure_collection(
164168
)
165169
else:
166170
raise QQLRuntimeError("Collection has no dense vector")
171+
return current_topology
167172
else:
168173
async with self._creation_lock:
169174
current_topology = await self._resolve_topology(name)
170175
if current_topology.exists:
171-
await self._ensure_collection(name, vector_size, current_topology, explicit_vector)
172-
return
176+
return await self._ensure_collection(name, vector_size, current_topology, explicit_vector)
173177

174178
await self._create_collection_and_wait(
175179
collection_name=name,
@@ -179,6 +183,7 @@ async def _ensure_collection(
179183
)
180184
},
181185
)
186+
return await self._resolve_topology(name)
182187

183188
async def _create_collection_and_wait(self, **kwargs: Any) -> None:
184189
collection_name = kwargs["collection_name"]
@@ -290,7 +295,7 @@ async def _execute_insert(self, node: InsertStmt) -> ExecutionResult:
290295
embedder = Embedder(model_name)
291296
vector = embedder.embed(node.values["text"])
292297

293-
await self._ensure_collection(
298+
topology = await self._ensure_collection(
294299
node.collection, len(vector), topology, node.dense_vector
295300
)
296301
point_vector = build_dense_point_vector(
@@ -351,22 +356,6 @@ async def _execute_insert_bulk(self, node: InsertBulkStmt) -> ExecutionResult:
351356
sparse_objs = [sparse_embedder.embed(vals["text"]) for vals in node.values_list]
352357

353358
first_dense_vector = dense_vectors[0] if dense_vectors else None
354-
points: list[PointStruct] = []
355-
for idx, vals in enumerate(node.values_list):
356-
point_id, payload = extract_point_id_and_payload(vals)
357-
dense_vector = dense_vectors[idx]
358-
sparse_obj = sparse_objs[idx]
359-
sparse_vector = SparseVector(
360-
indices=sparse_obj["indices"], values=sparse_obj["values"]
361-
)
362-
points.append(
363-
PointStruct(
364-
id=point_id,
365-
vector={dense_name: dense_vector, sparse_name: sparse_vector},
366-
payload=payload,
367-
)
368-
)
369-
370359
if not topology.exists:
371360
assert first_dense_vector is not None
372361
async with self._creation_lock:
@@ -385,6 +374,22 @@ async def _execute_insert_bulk(self, node: InsertBulkStmt) -> ExecutionResult:
385374
dense_name = current_topology.dense_using(node.dense_vector) or dense_name
386375
sparse_name = current_topology.sparse_using(node.sparse_vector)
387376

377+
points: list[PointStruct] = []
378+
for idx, vals in enumerate(node.values_list):
379+
point_id, payload = extract_point_id_and_payload(vals)
380+
dense_vector = dense_vectors[idx]
381+
sparse_obj = sparse_objs[idx]
382+
sparse_vector = SparseVector(
383+
indices=sparse_obj["indices"], values=sparse_obj["values"]
384+
)
385+
points.append(
386+
PointStruct(
387+
id=point_id,
388+
vector={dense_name: dense_vector, sparse_name: sparse_vector},
389+
payload=payload,
390+
)
391+
)
392+
388393
try:
389394
await self._client.upsert(
390395
collection_name=node.collection,
@@ -406,6 +411,10 @@ async def _execute_insert_bulk(self, node: InsertBulkStmt) -> ExecutionResult:
406411
vectors = [embedder.embed(vals["text"]) for vals in node.values_list]
407412

408413
first_vector = vectors[0] if vectors else None
414+
assert first_vector is not None
415+
topology = await self._ensure_collection(
416+
node.collection, len(first_vector), topology, node.dense_vector
417+
)
409418
points = []
410419
for idx, vals in enumerate(node.values_list):
411420
vector = vectors[idx]
@@ -420,11 +429,6 @@ async def _execute_insert_bulk(self, node: InsertBulkStmt) -> ExecutionResult:
420429
PointStruct(id=point_id, vector=point_vector, payload=payload)
421430
)
422431

423-
assert first_vector is not None
424-
await self._ensure_collection(
425-
node.collection, len(first_vector), topology, node.dense_vector
426-
)
427-
428432
try:
429433
await self._client.upsert(
430434
collection_name=node.collection,

src/qql/connection.py

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -198,22 +198,36 @@ def __exit__(self, exc_type: type[BaseException] | None, exc_val: BaseException
198198
if not self._queries:
199199
return
200200
results = self.connection.run_queries_batch(self._queries)
201-
for proxy, res in zip(self._proxies, results):
201+
for proxy, res in zip(self._proxies, results, strict=False):
202202
proxy._resolve(res)
203+
if len(results) != len(self._proxies):
204+
error = RuntimeError(
205+
"Batch result count mismatch: "
206+
f"expected {len(self._proxies)}, got {len(results)}"
207+
)
208+
for proxy in self._proxies[len(results):]:
209+
proxy._reject(error)
210+
raise error
203211

204212

205213
class OperationProxy:
206214
"""Proxy handle that resolves to an ExecutionResult after QQLBatch exits."""
207215

208216
def __init__(self) -> None:
209217
self._result: ExecutionResult | None = None
218+
self._exception: RuntimeError | None = None
210219

211220
def _resolve(self, result: ExecutionResult) -> None:
212221
self._result = result
213222

223+
def _reject(self, exception: RuntimeError) -> None:
224+
self._exception = exception
225+
214226
@property
215227
def result(self) -> ExecutionResult:
216228
"""The resolved ExecutionResult."""
229+
if self._exception is not None:
230+
raise self._exception
217231
if self._result is None:
218232
raise RuntimeError("Batch has not been executed yet.")
219233
return self._result

src/qql/parser.py

Lines changed: 10 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,8 @@ def __init__(self, tokens: list[Token]) -> None:
7272

7373
def parse(self) -> ASTNode:
7474
node = self._parse_single_statement()
75+
while self._peek().kind == TokenKind.SEMICOLON:
76+
self._advance()
7577
self._expect(TokenKind.EOF)
7678
return node
7779

@@ -1078,12 +1080,15 @@ def _parse_field_path(self) -> str:
10781080
f"Expected a field name, got '{tok.value}'", tok.pos
10791081
)
10801082

1081-
def _parse_literal(self) -> str | int | float | bool:
1082-
"""STRING | INTEGER | FLOAT | boolean"""
1083+
def _parse_literal(self) -> str | int | float | bool | None:
1084+
"""STRING | INTEGER | FLOAT | boolean | NULL"""
10831085
tok = self._peek()
10841086
if tok.kind == TokenKind.STRING:
10851087
self._advance()
10861088
return tok.value
1089+
if tok.kind == TokenKind.NULL:
1090+
self._advance()
1091+
return None
10871092
if tok.kind == TokenKind.INTEGER:
10881093
self._advance()
10891094
return int(tok.value)
@@ -1099,7 +1104,7 @@ def _parse_literal(self) -> str | int | float | bool:
10991104
self._advance()
11001105
return False
11011106
raise QQLSyntaxError(
1102-
f"Expected a literal value (string, integer, float, or boolean), got '{tok.value}'",
1107+
f"Expected a literal value (string, integer, float, boolean, or null), got '{tok.value}'",
11031108
tok.pos,
11041109
)
11051110

@@ -1116,10 +1121,10 @@ def _parse_number(self) -> int | float:
11161121
f"Expected a number, got '{tok.value}'", tok.pos
11171122
)
11181123

1119-
def _parse_literal_list(self) -> list[str | int | float | bool]:
1124+
def _parse_literal_list(self) -> list[str | int | float | bool | None]:
11201125
"""'(' literal { ',' literal } [','] ')' — used by IN / NOT IN."""
11211126
self._expect(TokenKind.LPAREN)
1122-
items: list[str | int | float | bool] = []
1127+
items: list[str | int | float | bool | None] = []
11231128
if self._peek().kind == TokenKind.RPAREN:
11241129
self._advance()
11251130
return items

src/qql/utils.py

Lines changed: 73 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -78,19 +78,65 @@ class SearchGroupByOptions:
7878

7979

8080
def render_parameterized_query(template: str, params: dict[str, Any]) -> str:
81-
query_str = template
82-
for k in sorted(params.keys(), key=len, reverse=True):
83-
val = params[k]
84-
placeholder = f":{k}"
85-
if isinstance(val, str):
86-
escaped_val = val.replace("'", "\\'")
87-
repr_val = f"'{escaped_val}'"
88-
elif isinstance(val, bool):
89-
repr_val = "true" if val else "false"
90-
else:
91-
repr_val = str(val)
92-
query_str = query_str.replace(placeholder, repr_val)
93-
return query_str
81+
rendered = []
82+
in_string = False
83+
quote_char = ""
84+
i = 0
85+
while i < len(template):
86+
ch = template[i]
87+
if in_string:
88+
rendered.append(ch)
89+
if ch == "\\" and i + 1 < len(template):
90+
rendered.append(template[i + 1])
91+
i += 2
92+
continue
93+
if ch == quote_char:
94+
in_string = False
95+
quote_char = ""
96+
i += 1
97+
continue
98+
99+
if ch in ("'", '"'):
100+
in_string = True
101+
quote_char = ch
102+
rendered.append(ch)
103+
i += 1
104+
continue
105+
106+
if ch == ":":
107+
name_start = i + 1
108+
name_end = name_start
109+
while name_end < len(template) and (
110+
template[name_end].isalnum() or template[name_end] == "_"
111+
):
112+
name_end += 1
113+
name = template[name_start:name_end]
114+
if name in params:
115+
rendered.append(_qql_literal(params[name]))
116+
i = name_end
117+
continue
118+
119+
rendered.append(ch)
120+
i += 1
121+
122+
return "".join(rendered)
123+
124+
125+
def _qql_literal(value: Any) -> str:
126+
if value is None:
127+
return "null"
128+
if isinstance(value, str):
129+
escaped = (
130+
value.replace("\\", "\\\\")
131+
.replace("'", "\\'")
132+
.replace("\n", "\\n")
133+
.replace("\t", "\\t")
134+
.replace("\r", "\\r")
135+
)
136+
return f"'{escaped}'"
137+
if isinstance(value, bool):
138+
return "true" if value else "false"
139+
return str(value)
94140

95141

96142
def collection_topology_kwargs(vectors: Any, sparse_vectors: Any) -> dict[str, Any]:
@@ -158,6 +204,14 @@ def group_batch_statements(statements: tuple[ASTNode, ...]) -> list[BatchGroup]:
158204
current_group: list[ASTNode] = []
159205

160206
for stmt in statements:
207+
if isinstance(stmt, SearchStmt) and stmt.group_by is not None:
208+
_append_batch_group(groups, current_type, current_collection, current_group)
209+
groups.append(BatchGroup("other", stmt.collection, [stmt]))
210+
current_type = None
211+
current_collection = None
212+
current_group = []
213+
continue
214+
161215
if isinstance(stmt, (SearchStmt, RecommendStmt)):
162216
coll = stmt.collection
163217
if current_type == "query" and current_collection == coll:
@@ -247,6 +301,12 @@ def build_qdrant_filter(expr: FilterExpr) -> Any:
247301
if isinstance(expr, NotExpr):
248302
return Filter(must_not=[build_qdrant_filter(expr.operand)])
249303
if isinstance(expr, CompareExpr):
304+
if expr.value is None:
305+
null_condition = IsNullCondition(is_null=PayloadField(key=expr.field))
306+
if expr.op == "=":
307+
return null_condition
308+
if expr.op == "!=":
309+
return Filter(must_not=[null_condition])
250310
if expr.op == "=":
251311
return FieldCondition(key=expr.field, match=MatchValue(value=expr.value))
252312
if expr.op == "!=":

0 commit comments

Comments
 (0)