langgraph-store-qdrant 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- langgraph/store/qdrant/__init__.py +6 -0
- langgraph/store/qdrant/_filters.py +273 -0
- langgraph/store/qdrant/_payload.py +160 -0
- langgraph/store/qdrant/_plans.py +96 -0
- langgraph/store/qdrant/aio.py +210 -0
- langgraph/store/qdrant/base.py +1047 -0
- langgraph/store/qdrant/py.typed +0 -0
- langgraph_store_qdrant-0.1.0.dist-info/METADATA +129 -0
- langgraph_store_qdrant-0.1.0.dist-info/RECORD +11 -0
- langgraph_store_qdrant-0.1.0.dist-info/WHEEL +4 -0
- langgraph_store_qdrant-0.1.0.dist-info/licenses/LICENSE +201 -0
|
@@ -0,0 +1,273 @@
|
|
|
1
|
+
"""Compiling store filters into Qdrant conditions.
|
|
2
|
+
|
|
3
|
+
Qdrant filters cannot tell an object from a scalar or a missing field from
|
|
4
|
+
`[]`, and match scalars against any element of an array. Each item's
|
|
5
|
+
`value_paths` markers (see `_payload.hash_value`) make those comparisons exact.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import itertools
|
|
11
|
+
import json
|
|
12
|
+
import math
|
|
13
|
+
import re
|
|
14
|
+
from collections.abc import Callable, Iterable, Mapping, Sequence
|
|
15
|
+
from contextlib import suppress
|
|
16
|
+
from dataclasses import dataclass
|
|
17
|
+
from datetime import datetime
|
|
18
|
+
from typing import Any
|
|
19
|
+
|
|
20
|
+
from langgraph.store.base import ListNamespacesOp, MatchCondition, SearchOp
|
|
21
|
+
from qdrant_client import models
|
|
22
|
+
|
|
23
|
+
from langgraph.store.qdrant._payload import (
|
|
24
|
+
MAX_INT,
|
|
25
|
+
MIN_INT,
|
|
26
|
+
dumps,
|
|
27
|
+
encode_ns,
|
|
28
|
+
hash_value,
|
|
29
|
+
number_key,
|
|
30
|
+
parse_dt,
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
Condition = models.Condition
|
|
34
|
+
|
|
35
|
+
# Qdrant compares ranges as f64, which is exact only below this magnitude.
|
|
36
|
+
_SAFE_FLOAT = 2**53
|
|
37
|
+
|
|
38
|
+
_SIMPLE_KEY = re.compile(r"[A-Za-z0-9_\-]+")
|
|
39
|
+
_NUMBER = re.compile(r"[+-]?(?:\d+\.?\d*|\.\d+)(?:[eE][+-]?\d+)?")
|
|
40
|
+
_INTEGER = re.compile(r"[+-]?\d+")
|
|
41
|
+
|
|
42
|
+
_AFFIX_FIELDS = {"prefix": "ns_prefixes", "suffix": "ns_suffixes"}
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def match(key: str, value: Any) -> models.FieldCondition:
|
|
46
|
+
return models.FieldCondition(key=key, match=models.MatchValue(value=value))
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def match_any(key: str, values: list[str]) -> models.FieldCondition:
|
|
50
|
+
return models.FieldCondition(key=key, match=models.MatchAny(any=values))
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def within(key: str, **bounds: Any) -> models.FieldCondition:
|
|
54
|
+
dated = any(isinstance(bound, datetime) for bound in bounds.values())
|
|
55
|
+
range_type = models.DatetimeRange if dated else models.Range
|
|
56
|
+
return models.FieldCondition(key=key, range=range_type(**bounds))
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def is_empty(key: str) -> models.IsEmptyCondition:
|
|
60
|
+
return models.IsEmptyCondition(is_empty=models.PayloadField(key=key))
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def every(*conditions: Condition) -> models.Filter:
|
|
64
|
+
return models.Filter(must=list(conditions))
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def either(*conditions: Condition) -> models.Filter:
|
|
68
|
+
return models.Filter(should=list(conditions))
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def negate(*conditions: Condition) -> models.Filter:
|
|
72
|
+
return models.Filter(must_not=list(conditions))
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def where(*conditions: Condition) -> models.Filter | None:
|
|
76
|
+
return every(*conditions) if conditions else None
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def not_expired(now: datetime) -> models.Filter:
|
|
80
|
+
return either(is_empty("expires_at"), within("expires_at", gt=now))
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def expired(now: datetime) -> models.Filter:
|
|
84
|
+
return every(within("expires_at", lte=now))
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
@dataclass(frozen=True)
|
|
88
|
+
class FieldPath:
|
|
89
|
+
parts: tuple[str, ...]
|
|
90
|
+
|
|
91
|
+
@property
|
|
92
|
+
def key(self) -> str:
|
|
93
|
+
if any(c in part for part in self.parts for c in '"\\'):
|
|
94
|
+
raise ValueError(
|
|
95
|
+
f"Filter keys cannot contain double quotes or backslashes: {self.parts}"
|
|
96
|
+
)
|
|
97
|
+
quoted = (p if _SIMPLE_KEY.fullmatch(p) else f'"{p}"' for p in self.parts)
|
|
98
|
+
return ".".join(["value", *quoted])
|
|
99
|
+
|
|
100
|
+
def child(self, name: str) -> FieldPath:
|
|
101
|
+
return FieldPath((*self.parts, name))
|
|
102
|
+
|
|
103
|
+
def has(self, marker: str) -> models.FieldCondition:
|
|
104
|
+
return match("value_paths", f"{marker}:{encode_ns(self.parts)}")
|
|
105
|
+
|
|
106
|
+
def equals_json(self, value: Any) -> models.FieldCondition:
|
|
107
|
+
"""Exact equality of an object or list, via its stored content hash."""
|
|
108
|
+
try:
|
|
109
|
+
normalized = json.loads(dumps(value))
|
|
110
|
+
except (TypeError, ValueError, RecursionError) as exc:
|
|
111
|
+
raise ValueError(f"Filter values must be valid JSON: {exc}") from exc
|
|
112
|
+
encoded = encode_ns(self.parts)
|
|
113
|
+
return match("value_paths", f"h:{encoded}:{hash_value(normalized).hex()}")
|
|
114
|
+
|
|
115
|
+
def equals_number(self, number: float) -> Condition:
|
|
116
|
+
"""Exact numeric equality, `1 == 1.0`: stored integers by value, stored
|
|
117
|
+
floats by their `n:` marker (Qdrant's own floats can be off by a ULP)."""
|
|
118
|
+
marker = match("value_paths", f"n:{encode_ns(self.parts)}:{number_key(number)}")
|
|
119
|
+
integral = isinstance(number, int) or number.is_integer()
|
|
120
|
+
if integral and MIN_INT <= int(number) <= MAX_INT:
|
|
121
|
+
return either(match(self.key, int(number)), marker)
|
|
122
|
+
return marker
|
|
123
|
+
|
|
124
|
+
def scalar(self, condition: Condition) -> list[Condition]:
|
|
125
|
+
# Qdrant matches scalar conditions against any element of an array.
|
|
126
|
+
return [condition, negate(self.has("a"))]
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def equals(path: FieldPath, value: Any, *, exact: bool = False) -> list[Condition]:
|
|
130
|
+
"""Conditions for `<value at path> == value`.
|
|
131
|
+
|
|
132
|
+
A plain object filter matches objects containing its keys. With `exact`
|
|
133
|
+
(`$eq`/`$ne`), the whole object must be equal.
|
|
134
|
+
"""
|
|
135
|
+
match value:
|
|
136
|
+
case Mapping() if exact:
|
|
137
|
+
return [path.equals_json(dict(value))]
|
|
138
|
+
case Mapping():
|
|
139
|
+
nested = (field_conditions(path.child(str(k)), v) for k, v in value.items())
|
|
140
|
+
return [path.has("o"), *itertools.chain.from_iterable(nested)]
|
|
141
|
+
case list() | tuple():
|
|
142
|
+
return [path.equals_json(list(value))]
|
|
143
|
+
case None: # missing or null, while `[]` is an array
|
|
144
|
+
return path.scalar(is_empty(path.key))
|
|
145
|
+
case bool() | str():
|
|
146
|
+
return path.scalar(match(path.key, value))
|
|
147
|
+
case float() if not math.isfinite(value):
|
|
148
|
+
raise ValueError(f"Filter values must be finite numbers, got {value!r}")
|
|
149
|
+
case int() | float():
|
|
150
|
+
return path.scalar(path.equals_number(value))
|
|
151
|
+
raise ValueError(f"Unsupported filter value type: {type(value).__name__}")
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def range_bound(op: str, value: Any) -> float | datetime:
|
|
155
|
+
match value:
|
|
156
|
+
case datetime():
|
|
157
|
+
return value
|
|
158
|
+
case bool():
|
|
159
|
+
pass
|
|
160
|
+
case float() if not math.isfinite(value):
|
|
161
|
+
pass
|
|
162
|
+
case int() | float() if _exact_bound(value):
|
|
163
|
+
return value
|
|
164
|
+
case int() | float():
|
|
165
|
+
raise ValueError(
|
|
166
|
+
f"Operator {op} cannot compare {value} exactly: Qdrant rounds it. "
|
|
167
|
+
"Integer bounds up to ±900719925474099 are always exact."
|
|
168
|
+
)
|
|
169
|
+
case str() if _INTEGER.fullmatch(text := value.strip()):
|
|
170
|
+
return range_bound(op, int(text))
|
|
171
|
+
case str() if _NUMBER.fullmatch(text := value.strip()):
|
|
172
|
+
return range_bound(op, float(text))
|
|
173
|
+
case str():
|
|
174
|
+
with suppress(ValueError):
|
|
175
|
+
return parse_dt(value.strip())
|
|
176
|
+
raise ValueError(
|
|
177
|
+
f"Operator {op} requires a finite number or an ISO-8601 datetime, got {value!r}"
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
def _exact_bound(number: float) -> bool:
|
|
182
|
+
if isinstance(number, float) and abs(number) >= 2**64:
|
|
183
|
+
return True # rounding cannot reach a stored integer
|
|
184
|
+
# A bound n is sent as "n.0", which Qdrant parses as (n * 10) / 10 in
|
|
185
|
+
# floating point: exact only if n * 10 is.
|
|
186
|
+
if abs(number) >= _SAFE_FLOAT:
|
|
187
|
+
return False
|
|
188
|
+
if isinstance(number, float) and not number.is_integer():
|
|
189
|
+
return True
|
|
190
|
+
return float(int(number) * 10) == int(number) * 10
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _comparison(bound: str) -> Callable[[FieldPath, Any], list[Condition]]:
|
|
194
|
+
def compare(path: FieldPath, operand: Any) -> list[Condition]:
|
|
195
|
+
limit = range_bound(f"${bound}", operand)
|
|
196
|
+
return path.scalar(within(path.key, **{bound: limit}))
|
|
197
|
+
|
|
198
|
+
return compare
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
OPERATORS: dict[str, Callable[[FieldPath, Any], list[Condition]]] = {
|
|
202
|
+
"$eq": lambda path, operand: equals(path, operand, exact=True),
|
|
203
|
+
"$ne": lambda path, operand: [negate(every(*equals(path, operand, exact=True)))],
|
|
204
|
+
**{f"${bound}": _comparison(bound) for bound in ("gt", "gte", "lt", "lte")},
|
|
205
|
+
}
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def field_conditions(path: FieldPath, value: Any) -> list[Condition]:
|
|
209
|
+
if not (isinstance(value, Mapping) and any(str(k).startswith("$") for k in value)):
|
|
210
|
+
return equals(path, value)
|
|
211
|
+
if unknown := [op for op in value if op not in OPERATORS]:
|
|
212
|
+
raise ValueError(f"Unsupported operator: {unknown[0]}")
|
|
213
|
+
return [c for op, operand in value.items() for c in OPERATORS[op](path, operand)]
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
def search_filter(
|
|
217
|
+
op: SearchOp, *, now: datetime, omit_expired: bool
|
|
218
|
+
) -> models.Filter | None:
|
|
219
|
+
# Stored keys are strings: values are JSON-normalized.
|
|
220
|
+
try:
|
|
221
|
+
value_conditions = [
|
|
222
|
+
field_conditions(FieldPath((str(key),)), value)
|
|
223
|
+
for key, value in (op.filter or {}).items()
|
|
224
|
+
]
|
|
225
|
+
except RecursionError:
|
|
226
|
+
raise ValueError("Filter is nested too deeply.") from None
|
|
227
|
+
return where(
|
|
228
|
+
*_present(
|
|
229
|
+
op.namespace_prefix, match("ns_prefixes", encode_ns(op.namespace_prefix))
|
|
230
|
+
),
|
|
231
|
+
*itertools.chain.from_iterable(value_conditions),
|
|
232
|
+
*_present(omit_expired, not_expired(now)),
|
|
233
|
+
)
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def _present(flag: object, condition: Condition) -> list[Condition]:
|
|
237
|
+
return [condition] if flag else []
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def _oriented(labels: Sequence[str], match_type: str) -> tuple[str, ...]:
|
|
241
|
+
"""Labels read from the end that `match_type` anchors on."""
|
|
242
|
+
return tuple(labels) if match_type == "prefix" else tuple(reversed(labels))
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
def _fixed_affix(condition: MatchCondition) -> tuple[str, ...]:
|
|
246
|
+
"""The part of the condition path before its first wildcard."""
|
|
247
|
+
oriented = _oriented(condition.path, condition.match_type)
|
|
248
|
+
fixed = itertools.takewhile(lambda label: label != "*", oriented)
|
|
249
|
+
return _oriented(tuple(fixed), condition.match_type)
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
def list_namespaces_filter(
|
|
253
|
+
op: ListNamespacesOp, *, now: datetime, omit_expired: bool
|
|
254
|
+
) -> models.Filter | None:
|
|
255
|
+
"""Server-side pre-filter. Wildcards are resolved by `namespace_matches`."""
|
|
256
|
+
conditions = [c for c in op.match_conditions or () if c.path]
|
|
257
|
+
return where(
|
|
258
|
+
*(within("ns_depth", gte=len(c.path)) for c in conditions),
|
|
259
|
+
*(
|
|
260
|
+
match(_AFFIX_FIELDS[c.match_type], encode_ns(fixed))
|
|
261
|
+
for c in conditions
|
|
262
|
+
if (fixed := _fixed_affix(c))
|
|
263
|
+
),
|
|
264
|
+
*_present(omit_expired, not_expired(now)),
|
|
265
|
+
)
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def namespace_matches(namespace: Iterable[str], condition: MatchCondition) -> bool:
|
|
269
|
+
path = _oriented(condition.path, condition.match_type)
|
|
270
|
+
labels = _oriented(tuple(namespace), condition.match_type)
|
|
271
|
+
return len(path) <= len(labels) and all(
|
|
272
|
+
p in ("*", label) for p, label in zip(path, labels, strict=False)
|
|
273
|
+
)
|
|
@@ -0,0 +1,160 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import hashlib
|
|
4
|
+
import json
|
|
5
|
+
import uuid
|
|
6
|
+
from collections.abc import Mapping, Sequence
|
|
7
|
+
from datetime import datetime, timedelta
|
|
8
|
+
from typing import Any, cast
|
|
9
|
+
|
|
10
|
+
MAX_VALUE_DEPTH = 32
|
|
11
|
+
"""qdrant-client's gRPC transport cannot encode objects nested any deeper."""
|
|
12
|
+
|
|
13
|
+
MIN_INT, MAX_INT = -(2**63), 2**63 - 1
|
|
14
|
+
|
|
15
|
+
_POINT_ID_NAMESPACE = uuid.UUID("5b2f8d0e-6c1a-4d0b-9a57-1f1f6a2c3e01")
|
|
16
|
+
|
|
17
|
+
dumps = json.JSONEncoder(
|
|
18
|
+
ensure_ascii=False, allow_nan=False, separators=(",", ":")
|
|
19
|
+
).encode
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def point_id(namespace: Sequence[str], key: str) -> str:
|
|
23
|
+
return str(uuid.uuid5(_POINT_ID_NAMESPACE, dumps([list(namespace), key])))
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def encode_ns(labels: Sequence[str]) -> str:
|
|
27
|
+
return dumps(list(labels))
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def parse_dt(value: str | datetime) -> datetime:
|
|
31
|
+
if isinstance(value, datetime):
|
|
32
|
+
return value
|
|
33
|
+
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def expires_at(updated_at: datetime, ttl: float | None) -> str | None:
|
|
37
|
+
return None if ttl is None else (updated_at + timedelta(minutes=ttl)).isoformat()
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def normalize_value(value: Mapping[str, Any]) -> dict[str, Any]:
|
|
41
|
+
"""The value exactly as stored and read back (a strict JSON round trip).
|
|
42
|
+
|
|
43
|
+
Rejects what Qdrant would store lossily or refuse.
|
|
44
|
+
"""
|
|
45
|
+
if not isinstance(value, Mapping):
|
|
46
|
+
raise TypeError(f"Store values must be mappings, got {type(value).__name__}")
|
|
47
|
+
try:
|
|
48
|
+
normalized = json.loads(dumps(dict(value)))
|
|
49
|
+
except ValueError as exc:
|
|
50
|
+
raise ValueError(f"Store values must be valid JSON: {exc}") from exc
|
|
51
|
+
except RecursionError:
|
|
52
|
+
raise ValueError(_TOO_DEEP) from None
|
|
53
|
+
_check_storable(normalized, depth=1)
|
|
54
|
+
return cast(dict[str, Any], normalized)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
_TOO_DEEP = f"Store values cannot be nested deeper than {MAX_VALUE_DEPTH} levels."
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def _check_storable(value: Any, depth: int) -> None:
|
|
61
|
+
match value:
|
|
62
|
+
case dict() | list() if depth > MAX_VALUE_DEPTH:
|
|
63
|
+
raise ValueError(_TOO_DEEP)
|
|
64
|
+
case dict():
|
|
65
|
+
for child in value.values():
|
|
66
|
+
_check_storable(child, depth + 1)
|
|
67
|
+
case list():
|
|
68
|
+
for child in value:
|
|
69
|
+
_check_storable(child, depth + 1)
|
|
70
|
+
case int() if not MIN_INT <= value <= MAX_INT:
|
|
71
|
+
raise ValueError(
|
|
72
|
+
f"Integer {value} is outside the 64-bit signed range Qdrant can "
|
|
73
|
+
"store and match exactly. Store it as a string instead."
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def number_key(number: float) -> str:
|
|
78
|
+
"""Canonical text of a number, the same for `1` and `1.0`."""
|
|
79
|
+
if isinstance(number, float) and not number.is_integer():
|
|
80
|
+
return repr(number)
|
|
81
|
+
return str(int(number))
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def _digest(data: bytes) -> bytes:
|
|
85
|
+
return hashlib.blake2b(data, digest_size=16).digest()
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def hash_value(
|
|
89
|
+
value: Any, path: tuple[str, ...] = (), markers: list[str] | None = None
|
|
90
|
+
) -> bytes:
|
|
91
|
+
"""Content hash where `1 == 1.0` and key order does not matter.
|
|
92
|
+
|
|
93
|
+
Appends `value_paths` markers to `markers`: `o:<path>` (object), `a:<path>`
|
|
94
|
+
(array) and `h:<path>:<hash>` for nested objects and arrays, and
|
|
95
|
+
`n:<path>:<number_key>` for floats.
|
|
96
|
+
"""
|
|
97
|
+
match value:
|
|
98
|
+
case dict():
|
|
99
|
+
kind = "o"
|
|
100
|
+
data = b",".join(
|
|
101
|
+
dumps(k).encode() + b":" + hash_value(value[k], (*path, k), markers)
|
|
102
|
+
for k in sorted(value)
|
|
103
|
+
)
|
|
104
|
+
case list():
|
|
105
|
+
kind = "a"
|
|
106
|
+
# Lists are compared as a whole, so their elements get no markers.
|
|
107
|
+
data = b",".join(hash_value(v) for v in value)
|
|
108
|
+
case float():
|
|
109
|
+
if markers is not None and path:
|
|
110
|
+
markers.append(f"n:{encode_ns(path)}:{number_key(value)}")
|
|
111
|
+
return _digest(b"s" + number_key(value).encode())
|
|
112
|
+
case _:
|
|
113
|
+
return _digest(b"s" + dumps(value).encode())
|
|
114
|
+
digest = _digest(kind.encode() + data)
|
|
115
|
+
if markers is not None and path:
|
|
116
|
+
encoded = encode_ns(path)
|
|
117
|
+
markers += [f"{kind}:{encoded}", f"h:{encoded}:{digest.hex()}"]
|
|
118
|
+
return digest
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def value_paths(value: dict[str, Any]) -> list[str]:
|
|
122
|
+
markers: list[str] = []
|
|
123
|
+
hash_value(value, markers=markers)
|
|
124
|
+
return markers
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def build_payload(
|
|
128
|
+
namespace: tuple[str, ...], key: str, value: dict[str, Any], ttl: float | None
|
|
129
|
+
) -> dict[str, Any]:
|
|
130
|
+
"""The payload without timestamps, which are set when writing."""
|
|
131
|
+
depths = range(1, len(namespace) + 1)
|
|
132
|
+
return {
|
|
133
|
+
"namespace": list(namespace),
|
|
134
|
+
"ns_path": encode_ns(namespace),
|
|
135
|
+
"ns_prefixes": [encode_ns(namespace[:d]) for d in depths],
|
|
136
|
+
"ns_suffixes": [encode_ns(namespace[-d:]) for d in depths],
|
|
137
|
+
"ns_depth": len(namespace),
|
|
138
|
+
"key": key,
|
|
139
|
+
"value": value,
|
|
140
|
+
# Qdrant can change floats by one unit in the last place and reorder
|
|
141
|
+
# keys, so items are read from this exact copy.
|
|
142
|
+
"value_json": dumps(value),
|
|
143
|
+
"value_paths": value_paths(value),
|
|
144
|
+
"ttl_minutes": ttl,
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def item_fields(payload: Mapping[str, Any]) -> dict[str, Any]:
|
|
149
|
+
return {
|
|
150
|
+
"namespace": tuple(payload["namespace"]),
|
|
151
|
+
"key": payload["key"],
|
|
152
|
+
"value": json.loads(payload["value_json"]),
|
|
153
|
+
"created_at": parse_dt(payload["created_at"]),
|
|
154
|
+
"updated_at": parse_dt(payload["updated_at"]),
|
|
155
|
+
}
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def is_expired(payload: Mapping[str, Any], now: datetime) -> bool:
|
|
159
|
+
expiry = payload.get("expires_at")
|
|
160
|
+
return expiry is not None and parse_dt(expiry) <= now
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""Plans: store logic written once as generators that yield requests.
|
|
2
|
+
|
|
3
|
+
`QdrantStore` and `AsyncQdrantStore` only execute the requests, synchronously
|
|
4
|
+
or with asyncio.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from collections.abc import Awaitable, Callable, Generator, Iterable
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from typing import Any, TypeVar, cast
|
|
12
|
+
|
|
13
|
+
import httpx
|
|
14
|
+
from qdrant_client.http.exceptions import ResponseHandlingException
|
|
15
|
+
|
|
16
|
+
T = TypeVar("T")
|
|
17
|
+
|
|
18
|
+
TRANSPORT_ATTEMPTS = 3
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True)
|
|
22
|
+
class Call:
|
|
23
|
+
method: str
|
|
24
|
+
args: tuple[Any, ...]
|
|
25
|
+
kwargs: dict[str, Any]
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass(frozen=True)
|
|
29
|
+
class EmbedDocuments:
|
|
30
|
+
texts: list[str]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
@dataclass(frozen=True)
|
|
34
|
+
class EmbedQueries:
|
|
35
|
+
queries: list[str]
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True)
|
|
39
|
+
class Gather:
|
|
40
|
+
"""Independent plans, which the async store runs concurrently."""
|
|
41
|
+
|
|
42
|
+
plans: list[Plan[Any]]
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
Request = Call | EmbedDocuments | EmbedQueries | Gather
|
|
46
|
+
Plan = Generator[Request, Any, T]
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def call(method: str, *args: Any, **kwargs: Any) -> Plan[Any]:
|
|
50
|
+
# Every request is idempotent, so transient transport failures are retried.
|
|
51
|
+
for attempt in range(1, TRANSPORT_ATTEMPTS + 1):
|
|
52
|
+
try:
|
|
53
|
+
return (yield Call(method, args, kwargs))
|
|
54
|
+
except ResponseHandlingException as exc:
|
|
55
|
+
if attempt == TRANSPORT_ATTEMPTS or not is_transient(exc):
|
|
56
|
+
raise
|
|
57
|
+
raise AssertionError("unreachable")
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def gather(plans: Iterable[Plan[Any]]) -> Plan[list[Any]]:
|
|
61
|
+
return (yield Gather(list(plans)))
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def is_transient(exc: ResponseHandlingException) -> bool:
|
|
65
|
+
"""A transport failure without a response, e.g. a dropped connection."""
|
|
66
|
+
return isinstance(exc.source, (httpx.TransportError, OSError)) and not isinstance(
|
|
67
|
+
exc.source, httpx.TimeoutException
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def run(plan: Plan[T], execute: Callable[[Request], Any]) -> T:
|
|
72
|
+
try:
|
|
73
|
+
request = next(plan)
|
|
74
|
+
while True:
|
|
75
|
+
try:
|
|
76
|
+
result = execute(request)
|
|
77
|
+
except BaseException as exc:
|
|
78
|
+
request = plan.throw(exc)
|
|
79
|
+
else:
|
|
80
|
+
request = plan.send(result)
|
|
81
|
+
except StopIteration as stop:
|
|
82
|
+
return cast(T, stop.value)
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
async def arun(plan: Plan[T], execute: Callable[[Request], Awaitable[Any]]) -> T:
|
|
86
|
+
try:
|
|
87
|
+
request = next(plan)
|
|
88
|
+
while True:
|
|
89
|
+
try:
|
|
90
|
+
result = await execute(request)
|
|
91
|
+
except BaseException as exc:
|
|
92
|
+
request = plan.throw(exc)
|
|
93
|
+
else:
|
|
94
|
+
request = plan.send(result)
|
|
95
|
+
except StopIteration as stop:
|
|
96
|
+
return cast(T, stop.value)
|