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.
@@ -0,0 +1,6 @@
1
+ """Qdrant-backed long-term memory store for LangGraph."""
2
+
3
+ from langgraph.store.qdrant.aio import AsyncQdrantStore
4
+ from langgraph.store.qdrant.base import QdrantIndexConfig, QdrantStore
5
+
6
+ __all__ = ["AsyncQdrantStore", "QdrantIndexConfig", "QdrantStore"]
@@ -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)