cairn-db-client 0.3.2__tar.gz

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,5 @@
1
+ __pycache__/
2
+ *.pyc
3
+ .venv/
4
+ dist/
5
+ .pytest_cache/
@@ -0,0 +1,102 @@
1
+ Metadata-Version: 2.5
2
+ Name: cairn-db-client
3
+ Version: 0.3.2
4
+ Summary: Client for Cairn, the hybrid search database where a deletion is final
5
+ Project-URL: Homepage, https://github.com/Cairn-DB/cairn
6
+ Project-URL: Repository, https://github.com/Cairn-DB/cairn
7
+ Project-URL: Issues, https://github.com/Cairn-DB/cairn/issues
8
+ Project-URL: Changelog, https://github.com/Cairn-DB/cairn/blob/main/CHANGELOG.md
9
+ License-Expression: Apache-2.0
10
+ Keywords: cairn,database,hybrid-search,rag,vector-search
11
+ Classifier: Programming Language :: Python :: 3
12
+ Classifier: Typing :: Typed
13
+ Requires-Python: >=3.9
14
+ Requires-Dist: httpx>=0.24
15
+ Provides-Extra: test
16
+ Requires-Dist: pytest>=7; extra == 'test'
17
+ Description-Content-Type: text/markdown
18
+
19
+ # cairn-db-client
20
+
21
+ Python client for [Cairn](https://github.com/Cairn-DB/cairn), the hybrid search database
22
+ where a deletion is final. It comes in two flavours, sync (`Client`) and async
23
+ (`AsyncClient`), and is typed. Its only dependency is `httpx`.
24
+
25
+ ```bash
26
+ pip install cairn-db-client
27
+ ```
28
+
29
+ ```python
30
+ from cairn_db import Client, eq, range_, and_
31
+
32
+ db = Client("http://localhost:7200", api_key=KEY)
33
+
34
+ # Your own ids; chunks carry their parent in an ordinary field ("parent" by default).
35
+ db.upsert([
36
+ {"id": "report-9#0", "parent": "report-9", "text": "...", "embedding": [...], "lang": "en"},
37
+ {"id": "report-9#1", "parent": "report-9", "text": "...", "embedding": [...], "lang": "en"},
38
+ ])
39
+
40
+ hits = db.search(
41
+ k=10,
42
+ vector={"field": "embedding", "values": query_vector},
43
+ text="nuclear energy", text_field="text",
44
+ filter=and_(eq("lang", "en"), range_("year", gte=2020)),
45
+ )
46
+ for h in hits:
47
+ print(h.id, h.score, h.document["text"])
48
+
49
+ db.delete(parent="report-9") # the document and all its chunks
50
+ db.delete(filter=eq("source", "crawler")) # everything that matches, now
51
+ db.delete(ids=["report-7#0"])
52
+ ```
53
+
54
+ `AsyncClient` has the same methods, as coroutines (`async with AsyncClient(...) as db:`).
55
+
56
+ ## Read-your-writes and takedowns
57
+
58
+ The client keeps the consistency token of its writes and takedowns, and sends it with every
59
+ read. It therefore reads what it wrote, and never reads back what it deleted, through any
60
+ node. To give the same guarantee to another service, pass it `db.token`:
61
+ `Client(..., token=token)`, or `db.observe(token)` on an existing client.
62
+
63
+ ## Collections
64
+
65
+ Without `collection()`, calls act on `default`, the collection defined at startup.
66
+
67
+ ```python
68
+ admin = Client(url, api_key=ADMIN_KEY)
69
+ admin.create_collection("notes", {"fields": [
70
+ {"name": "embedding", "kind": {"Vector": {"dims": 384, "metric": "Cosine"}}},
71
+ {"name": "body", "kind": "Text"},
72
+ ]}, shards=4)
73
+ notes = db.collection("notes") # the same calls, on "notes"
74
+ notes.upsert([{"id": "n1", "embedding": [...], "body": "..."}])
75
+ admin.drop_collection("notes") # deletes its data on every node
76
+ ```
77
+
78
+ `create_collection` and `drop_collection` need an admin key. `list_collections` needs a read key.
79
+
80
+ ## Tenants
81
+
82
+ ```python
83
+ acme = db.with_tenant("acme") # an unscoped key acting for tenant "acme"
84
+ acme.upsert([{"id": "note-1", "text": "...", "embedding": [...]}])
85
+ db.forget_tenant("acme") # erase everything of "acme"
86
+ ```
87
+
88
+ A key scoped to a tenant (`cairn-server keygen app read,write --tenant acme`) needs nothing
89
+ else: every call stays inside its tenant. Ids belong to their tenant.
90
+
91
+ ## Errors and retries
92
+
93
+ - `InvalidInputError` (400), `AuthenticationError` (401), `ForbiddenError` (403),
94
+ `UnavailableError` (503 or network), all subclasses of `CairnError` with a `status`.
95
+ - `get` returns `None` for a missing document.
96
+ - 502, 503, 504 and network errors are retried (`retries=3` by default) with backoff, on the
97
+ next address when `url` is a list.
98
+
99
+ ## Tests
100
+
101
+ `PYTHONPATH=src python -m pytest tests` runs the unit tests. `../test-live.sh` also runs the
102
+ live tests against a fresh local node.
@@ -0,0 +1,84 @@
1
+ # cairn-db-client
2
+
3
+ Python client for [Cairn](https://github.com/Cairn-DB/cairn), the hybrid search database
4
+ where a deletion is final. It comes in two flavours, sync (`Client`) and async
5
+ (`AsyncClient`), and is typed. Its only dependency is `httpx`.
6
+
7
+ ```bash
8
+ pip install cairn-db-client
9
+ ```
10
+
11
+ ```python
12
+ from cairn_db import Client, eq, range_, and_
13
+
14
+ db = Client("http://localhost:7200", api_key=KEY)
15
+
16
+ # Your own ids; chunks carry their parent in an ordinary field ("parent" by default).
17
+ db.upsert([
18
+ {"id": "report-9#0", "parent": "report-9", "text": "...", "embedding": [...], "lang": "en"},
19
+ {"id": "report-9#1", "parent": "report-9", "text": "...", "embedding": [...], "lang": "en"},
20
+ ])
21
+
22
+ hits = db.search(
23
+ k=10,
24
+ vector={"field": "embedding", "values": query_vector},
25
+ text="nuclear energy", text_field="text",
26
+ filter=and_(eq("lang", "en"), range_("year", gte=2020)),
27
+ )
28
+ for h in hits:
29
+ print(h.id, h.score, h.document["text"])
30
+
31
+ db.delete(parent="report-9") # the document and all its chunks
32
+ db.delete(filter=eq("source", "crawler")) # everything that matches, now
33
+ db.delete(ids=["report-7#0"])
34
+ ```
35
+
36
+ `AsyncClient` has the same methods, as coroutines (`async with AsyncClient(...) as db:`).
37
+
38
+ ## Read-your-writes and takedowns
39
+
40
+ The client keeps the consistency token of its writes and takedowns, and sends it with every
41
+ read. It therefore reads what it wrote, and never reads back what it deleted, through any
42
+ node. To give the same guarantee to another service, pass it `db.token`:
43
+ `Client(..., token=token)`, or `db.observe(token)` on an existing client.
44
+
45
+ ## Collections
46
+
47
+ Without `collection()`, calls act on `default`, the collection defined at startup.
48
+
49
+ ```python
50
+ admin = Client(url, api_key=ADMIN_KEY)
51
+ admin.create_collection("notes", {"fields": [
52
+ {"name": "embedding", "kind": {"Vector": {"dims": 384, "metric": "Cosine"}}},
53
+ {"name": "body", "kind": "Text"},
54
+ ]}, shards=4)
55
+ notes = db.collection("notes") # the same calls, on "notes"
56
+ notes.upsert([{"id": "n1", "embedding": [...], "body": "..."}])
57
+ admin.drop_collection("notes") # deletes its data on every node
58
+ ```
59
+
60
+ `create_collection` and `drop_collection` need an admin key. `list_collections` needs a read key.
61
+
62
+ ## Tenants
63
+
64
+ ```python
65
+ acme = db.with_tenant("acme") # an unscoped key acting for tenant "acme"
66
+ acme.upsert([{"id": "note-1", "text": "...", "embedding": [...]}])
67
+ db.forget_tenant("acme") # erase everything of "acme"
68
+ ```
69
+
70
+ A key scoped to a tenant (`cairn-server keygen app read,write --tenant acme`) needs nothing
71
+ else: every call stays inside its tenant. Ids belong to their tenant.
72
+
73
+ ## Errors and retries
74
+
75
+ - `InvalidInputError` (400), `AuthenticationError` (401), `ForbiddenError` (403),
76
+ `UnavailableError` (503 or network), all subclasses of `CairnError` with a `status`.
77
+ - `get` returns `None` for a missing document.
78
+ - 502, 503, 504 and network errors are retried (`retries=3` by default) with backoff, on the
79
+ next address when `url` is a list.
80
+
81
+ ## Tests
82
+
83
+ `PYTHONPATH=src python -m pytest tests` runs the unit tests. `../test-live.sh` also runs the
84
+ live tests against a fresh local node.
@@ -0,0 +1,30 @@
1
+ [build-system]
2
+ requires = ["hatchling>=1.24"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "cairn-db-client"
7
+ version = "0.3.2"
8
+ description = "Client for Cairn, the hybrid search database where a deletion is final"
9
+ readme = "README.md"
10
+ license = "Apache-2.0"
11
+ requires-python = ">=3.9"
12
+ dependencies = ["httpx>=0.24"]
13
+ classifiers = ["Typing :: Typed", "Programming Language :: Python :: 3"]
14
+
15
+ keywords = ["cairn", "vector-search", "hybrid-search", "rag", "database"]
16
+
17
+ [project.urls]
18
+ Homepage = "https://github.com/Cairn-DB/cairn"
19
+ Repository = "https://github.com/Cairn-DB/cairn"
20
+ Issues = "https://github.com/Cairn-DB/cairn/issues"
21
+ Changelog = "https://github.com/Cairn-DB/cairn/blob/main/CHANGELOG.md"
22
+
23
+ [project.optional-dependencies]
24
+ test = ["pytest>=7"]
25
+
26
+ [tool.hatch.build.targets.wheel]
27
+ packages = ["src/cairn_db"]
28
+
29
+ [tool.pytest.ini_options]
30
+ testpaths = ["tests"]
@@ -0,0 +1,37 @@
1
+ """``cairn-db-client`` (imported as ``cairn_db``): the Cairn HTTP API from Python, sync and async (ADR 0031)."""
2
+
3
+ from ._core import (
4
+ AuthenticationError,
5
+ CairnError,
6
+ Consistency,
7
+ DeleteResult,
8
+ Document,
9
+ Filter,
10
+ ForbiddenError,
11
+ Hit,
12
+ Id,
13
+ InvalidInputError,
14
+ LegHit,
15
+ NotFoundError,
16
+ UnavailableError,
17
+ VectorLeg,
18
+ WriteResult,
19
+ and_,
20
+ eq,
21
+ in_,
22
+ is_null,
23
+ merge_tokens,
24
+ not_,
25
+ or_,
26
+ range_,
27
+ )
28
+ from .client import AsyncClient, Client
29
+
30
+ __version__ = "0.3.2"
31
+
32
+ __all__ = [
33
+ "AsyncClient", "AuthenticationError", "CairnError", "Client", "Consistency", "DeleteResult",
34
+ "Document", "Filter", "ForbiddenError", "Hit", "Id", "InvalidInputError", "LegHit",
35
+ "NotFoundError", "UnavailableError", "VectorLeg", "WriteResult", "and_", "eq", "in_",
36
+ "is_null", "merge_tokens", "not_", "or_", "range_",
37
+ ]
@@ -0,0 +1,283 @@
1
+ """Request building, tokens, results and errors shared by the sync and async clients."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import re
6
+ import urllib.parse
7
+ from dataclasses import dataclass, field
8
+ from typing import Any, Literal, Mapping, Optional, Sequence, TypedDict, Union
9
+
10
+ Id = Union[str, int]
11
+ """A document id: your own string, or an unsigned integer below 2^63."""
12
+
13
+ Document = dict[str, Any]
14
+ """A document: an ``id`` plus fields named as in the schema."""
15
+
16
+ Filter = dict[str, Any]
17
+ """A filter: ``and``, ``or``, ``not``, or a condition on one field (see the helpers)."""
18
+
19
+ Consistency = Literal["linearizable", "read_your_writes", "stale"]
20
+
21
+
22
+ class VectorLeg(TypedDict, total=False):
23
+ field: str
24
+ values: list[float]
25
+ ef: int
26
+
27
+
28
+ # ---------------------------------------------------------------- filters
29
+
30
+
31
+ def eq(field: str, value: Any) -> Filter:
32
+ """``field`` equals ``value`` (on a Set field: contains it)."""
33
+ return {"field": field, "eq": value}
34
+
35
+
36
+ def in_(field: str, values: Sequence[Any]) -> Filter:
37
+ """``field`` equals one of ``values`` (on a Set field: contains one of them)."""
38
+ return {"field": field, "in": list(values)}
39
+
40
+
41
+ def range_(field: str, *, gt: Any = None, gte: Any = None, lt: Any = None, lte: Any = None) -> Filter:
42
+ """A numeric or date range."""
43
+ f: Filter = {"field": field}
44
+ for k, v in (("gt", gt), ("gte", gte), ("lt", lt), ("lte", lte)):
45
+ if v is not None:
46
+ f[k] = v
47
+ return f
48
+
49
+
50
+ def is_null(field: str) -> Filter:
51
+ """``field`` has no value."""
52
+ return {"field": field, "is_null": True}
53
+
54
+
55
+ def and_(*filters: Filter) -> Filter:
56
+ return {"and": list(filters)}
57
+
58
+
59
+ def or_(*filters: Filter) -> Filter:
60
+ return {"or": list(filters)}
61
+
62
+
63
+ def not_(f: Filter) -> Filter:
64
+ return {"not": f}
65
+
66
+
67
+ # ---------------------------------------------------------------- errors
68
+
69
+
70
+ class CairnError(Exception):
71
+ """Any error from the API or the network. ``status`` is 0 for a network error."""
72
+
73
+ def __init__(self, message: str, status: int) -> None:
74
+ super().__init__(message)
75
+ self.message = message
76
+ self.status = status
77
+
78
+
79
+ class AuthenticationError(CairnError):
80
+ """401: missing or invalid API key."""
81
+
82
+
83
+ class ForbiddenError(CairnError):
84
+ """403: the key lacks the role, or cannot act for that tenant."""
85
+
86
+
87
+ class NotFoundError(CairnError):
88
+ """404."""
89
+
90
+
91
+ class InvalidInputError(CairnError):
92
+ """400: invalid input (unknown field, wrong type, bad filter...)."""
93
+
94
+
95
+ class UnavailableError(CairnError):
96
+ """503, or a network failure, after the retries."""
97
+
98
+
99
+ def error_for(status: int, message: str) -> CairnError:
100
+ cls = {
101
+ 400: InvalidInputError,
102
+ 401: AuthenticationError,
103
+ 403: ForbiddenError,
104
+ 404: NotFoundError,
105
+ 0: UnavailableError,
106
+ 502: UnavailableError,
107
+ 503: UnavailableError,
108
+ 504: UnavailableError,
109
+ }.get(status, CairnError)
110
+ return cls(message, status)
111
+
112
+
113
+ # ---------------------------------------------------------------- results
114
+
115
+
116
+ @dataclass
117
+ class WriteResult:
118
+ """Documents written (or ids given to a takedown by id), and the consistency token."""
119
+
120
+ count: int
121
+ token: str
122
+
123
+
124
+ @dataclass
125
+ class DeleteResult:
126
+ """``deleted``: documents removed (by filter, parent or tenant); ``count``: ids given."""
127
+
128
+ token: str
129
+ deleted: Optional[int] = None
130
+ count: Optional[int] = None
131
+
132
+
133
+ @dataclass
134
+ class LegHit:
135
+ rank: int
136
+ score: float
137
+
138
+
139
+ @dataclass
140
+ class Hit:
141
+ id: Id
142
+ score: float
143
+ legs: list[Optional[LegHit]] = field(default_factory=list)
144
+ document: Optional[Document] = None
145
+ tenant: Optional[str] = None
146
+ """The hit's tenant, shown to unscoped keys only."""
147
+ group: Any = None
148
+ """The hit's value of the ``group_by`` field."""
149
+
150
+ @staticmethod
151
+ def from_json(h: Mapping[str, Any]) -> "Hit":
152
+ return Hit(
153
+ id=h["id"],
154
+ score=h["score"],
155
+ legs=[LegHit(**leg) if leg else None for leg in h.get("legs", [])],
156
+ document=h.get("document"),
157
+ tenant=h.get("_tenant"),
158
+ group=h.get("group"),
159
+ )
160
+
161
+
162
+ # ---------------------------------------------------------------- tokens and requests
163
+
164
+
165
+ def merge_tokens(*tokens: Optional[str]) -> str:
166
+ """Per-shard maximum of consistency tokens (``shard.index,...``)."""
167
+ best: dict[int, int] = {}
168
+ for t in tokens:
169
+ for part in (t or "").split(","):
170
+ if not part:
171
+ continue
172
+ m = re.fullmatch(r"(\d+)\.(\d+)", part)
173
+ if not m:
174
+ raise InvalidInputError(f"bad consistency token {part!r}", 400)
175
+ s, i = int(m.group(1)), int(m.group(2))
176
+ best[s] = max(best.get(s, 0), i)
177
+ return ",".join(f"{s}.{i}" for s, i in sorted(best.items()))
178
+
179
+
180
+ class TokenBox:
181
+ """Token state shared by a client and its tenant views."""
182
+
183
+ def __init__(self, value: str = "") -> None:
184
+ self.value = value
185
+
186
+ def observe(self, token: str) -> None:
187
+ self.value = merge_tokens(self.value, token)
188
+
189
+
190
+ def collection_base(collection: Optional[str]) -> str:
191
+ """Path prefix of a collection's routes."""
192
+ if not collection or collection == "default":
193
+ return "/v1"
194
+ return "/v1/collections/" + urllib.parse.quote(collection, safe="")
195
+
196
+
197
+ def doc_path(id: Id, after: str, consistency: Optional[str], base: str = "/v1") -> str:
198
+ """The path of a document: digits are an integer id, so a text id of digits says so."""
199
+ params: list[tuple[str, str]] = []
200
+ if isinstance(id, bool) or not isinstance(id, (str, int)):
201
+ raise InvalidInputError(f"an id is a string or an integer, not {id!r}", 400)
202
+ if isinstance(id, int):
203
+ if id < 0:
204
+ raise InvalidInputError(f"an integer id is non-negative: {id}", 400)
205
+ path = f"{base}/documents/{id}"
206
+ else:
207
+ path = f"{base}/documents/" + urllib.parse.quote(id, safe="")
208
+ if id.isdigit():
209
+ params.append(("id_type", "text"))
210
+ if after:
211
+ params.append(("after", after))
212
+ if consistency:
213
+ params.append(("consistency", consistency))
214
+ return path + ("?" + urllib.parse.urlencode(params) if params else "")
215
+
216
+
217
+ def search_body(
218
+ *,
219
+ k: Optional[int],
220
+ vector: Union[VectorLeg, Sequence[VectorLeg], None],
221
+ text: Optional[str],
222
+ text_field: Optional[str],
223
+ all_terms: bool,
224
+ filter: Optional[Filter],
225
+ fusion: Optional[Mapping[str, Any]],
226
+ oversample: Optional[int],
227
+ with_documents: Optional[bool],
228
+ consistency: Optional[str],
229
+ after: str,
230
+ group_by: Optional[str] = None,
231
+ ) -> dict[str, Any]:
232
+ body: dict[str, Any] = {}
233
+ if k is not None:
234
+ body["k"] = k
235
+ if isinstance(vector, Mapping):
236
+ body["vector"] = dict(vector)
237
+ elif vector:
238
+ body["vectors"] = [dict(v) for v in vector]
239
+ if text is not None:
240
+ if not text_field:
241
+ raise InvalidInputError("a text query needs text_field", 400)
242
+ body["text"] = {"field": text_field, "query": text, "all_terms": all_terms}
243
+ if filter is not None:
244
+ body["filter"] = filter
245
+ if fusion is not None:
246
+ body["fusion"] = dict(fusion)
247
+ if oversample is not None:
248
+ body["oversample"] = oversample
249
+ if with_documents is not None:
250
+ body["with_documents"] = with_documents
251
+ if consistency:
252
+ body["consistency"] = consistency
253
+ if group_by:
254
+ body["group_by"] = group_by
255
+ if after:
256
+ body["after"] = after
257
+ return body
258
+
259
+
260
+ def delete_body(
261
+ *, ids: Optional[Sequence[Id]], filter: Optional[Filter], parent: Optional[Id], parent_field: str, after: str
262
+ ) -> dict[str, Any]:
263
+ if parent is not None:
264
+ if ids is not None or filter is not None:
265
+ raise InvalidInputError("parent cannot be combined with ids or filter", 400)
266
+ body: dict[str, Any] = {"filter": eq(parent_field, parent)}
267
+ elif ids is not None:
268
+ if not ids:
269
+ raise InvalidInputError("no ids", 400)
270
+ body = {"ids": list(ids)}
271
+ if filter is not None:
272
+ body["filter"] = filter
273
+ elif filter is not None:
274
+ body = {"filter": filter}
275
+ else:
276
+ raise InvalidInputError("give ids, filter or parent", 400)
277
+ if after:
278
+ body["after"] = after
279
+ return body
280
+
281
+
282
+ def delete_result(r: Mapping[str, Any]) -> DeleteResult:
283
+ return DeleteResult(token=r["consistency_token"], deleted=r.get("deleted"), count=r.get("count"))
@@ -0,0 +1,403 @@
1
+ """Sync and async clients for the Cairn HTTP API (ADR 0031)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import time
7
+ import urllib.parse
8
+ from typing import Any, Mapping, Optional, Sequence, Union
9
+
10
+ import httpx
11
+
12
+ from ._core import (
13
+ Consistency,
14
+ DeleteResult,
15
+ Document,
16
+ Filter,
17
+ Hit,
18
+ Id,
19
+ NotFoundError,
20
+ TokenBox,
21
+ UnavailableError,
22
+ VectorLeg,
23
+ WriteResult,
24
+ collection_base,
25
+ delete_body,
26
+ delete_result,
27
+ doc_path,
28
+ error_for,
29
+ merge_tokens,
30
+ search_body,
31
+ )
32
+
33
+
34
+ class _Base:
35
+ def __init__(
36
+ self,
37
+ url: Union[str, Sequence[str]],
38
+ api_key: Optional[str] = None,
39
+ *,
40
+ tenant: Optional[str] = None,
41
+ parent_field: str = "parent",
42
+ timeout: float = 30.0,
43
+ retries: int = 3,
44
+ token: Optional[str] = None,
45
+ collection: Optional[str] = None,
46
+ _box: Optional[TokenBox] = None,
47
+ ) -> None:
48
+ urls = [url] if isinstance(url, str) else list(url)
49
+ if not urls:
50
+ raise ValueError("no url")
51
+ self._urls = [u.rstrip("/") for u in urls]
52
+ self._api_key = api_key
53
+ self._tenant = tenant
54
+ self._parent_field = parent_field
55
+ self._timeout = timeout
56
+ self._retries = retries
57
+ self._box = _box or TokenBox(merge_tokens(token) if token else "")
58
+ self._collection = collection
59
+ self._base = collection_base(collection)
60
+ self._next = 0
61
+
62
+ @property
63
+ def token(self) -> str:
64
+ """The consistency token covering every write and takedown of this client and its
65
+ tenant views. Hand it to another service so that its reads reflect them."""
66
+ return self._box.value
67
+
68
+ def observe(self, token: str) -> None:
69
+ """Adds a token received from elsewhere: later reads reflect those writes too."""
70
+ self._box.observe(token)
71
+
72
+ def _headers(self) -> dict[str, str]:
73
+ h = {"content-type": "application/json"}
74
+ if self._api_key:
75
+ h["authorization"] = f"Bearer {self._api_key}"
76
+ if self._tenant:
77
+ h["cairn-tenant"] = self._tenant
78
+ return h
79
+
80
+ def _view_args(self, **changes: Any) -> dict[str, Any]:
81
+ args: dict[str, Any] = dict(
82
+ url=self._urls,
83
+ api_key=self._api_key,
84
+ tenant=self._tenant,
85
+ parent_field=self._parent_field,
86
+ timeout=self._timeout,
87
+ retries=self._retries,
88
+ collection=self._collection,
89
+ _box=self._box,
90
+ )
91
+ args.update(changes)
92
+ return args
93
+
94
+ def _outcome(self, res: Optional[httpx.Response], err: Optional[Exception], url: str) -> Any:
95
+ """The decoded answer, or the error to raise (retry when it is an UnavailableError)."""
96
+ if res is None:
97
+ return UnavailableError(f"{url}: {err}", 0)
98
+ try:
99
+ data = res.json() if res.content else None
100
+ except ValueError:
101
+ data = None
102
+ if res.is_success:
103
+ return data
104
+ message = (data or {}).get("error") if isinstance(data, dict) else None
105
+ return error_for(res.status_code, message or res.text or res.reason_phrase)
106
+
107
+
108
+ class Client(_Base):
109
+ """Blocking client.
110
+
111
+ >>> db = Client("http://localhost:7200", api_key=KEY)
112
+ >>> db.upsert([{"id": "doc-1#0", "parent": "doc-1", "text": "...", "embedding": [...]}])
113
+ >>> hits = db.search(text="nuclear", text_field="text", filter=eq("lang", "en"))
114
+ >>> db.delete(parent="doc-1") # the document and all its chunks
115
+ >>> db.with_tenant("acme").get("doc-1") # an unscoped key acting for one tenant
116
+ >>> db.collection("notes").search(...) # another collection, same calls
117
+ >>> db.forget_tenant("acme") # erase a tenant
118
+
119
+ Reads pass the token of this client's writes and takedowns, so it reads its own writes and
120
+ never reads back what it deleted, through any node.
121
+ """
122
+
123
+ def __init__(self, url: Union[str, Sequence[str]], api_key: Optional[str] = None, **kw: Any) -> None:
124
+ super().__init__(url, api_key, **kw)
125
+ self._http = httpx.Client(timeout=self._timeout)
126
+
127
+ def close(self) -> None:
128
+ self._http.close()
129
+
130
+ def __enter__(self) -> "Client":
131
+ return self
132
+
133
+ def __exit__(self, *exc: object) -> None:
134
+ self.close()
135
+
136
+ def with_tenant(self, tenant: str) -> "Client":
137
+ """A view acting for ``tenant`` (unscoped keys), sharing this client's token."""
138
+ return Client(**self._view_args(tenant=tenant))
139
+
140
+ def collection(self, name: str) -> "Client":
141
+ """A view acting on collection ``name``, sharing this client's token and tenant."""
142
+ return Client(**self._view_args(collection=name))
143
+
144
+ def create_collection(
145
+ self,
146
+ name: str,
147
+ schema: Mapping[str, Any],
148
+ *,
149
+ shards: Optional[int] = None,
150
+ expires_field: Optional[str] = None,
151
+ ) -> dict[str, Any]:
152
+ """Creates a collection (admin key); returns once it is ready on every node. With
153
+ ``expires_field`` (a field holding each document's expiry, Unix milliseconds), expired
154
+ documents are hidden at once and then deleted."""
155
+ body: dict[str, Any] = {"name": name, "schema": dict(schema)}
156
+ if shards is not None:
157
+ body["shards"] = shards
158
+ if expires_field:
159
+ body["expires_field"] = expires_field
160
+ return self._call("POST", "/v1/collections", body)
161
+
162
+ def list_collections(self) -> list[dict[str, Any]]:
163
+ """The live collections, ``default`` first."""
164
+ return self._call("GET", "/v1/collections")["collections"]
165
+
166
+ def drop_collection(self, name: str) -> None:
167
+ """Drops a collection and deletes its data on every node (admin key)."""
168
+ self._call("DELETE", "/v1/collections/" + urllib.parse.quote(name, safe=""))
169
+
170
+ def _call(self, method: str, path: str, body: Any = None) -> Any:
171
+ last: Exception = UnavailableError("no attempt", 0)
172
+ for attempt in range(self._retries + 1):
173
+ if attempt:
174
+ time.sleep(0.1 * 2 ** (attempt - 1))
175
+ url = self._urls[self._next % len(self._urls)] + path
176
+ res, err = None, None
177
+ try:
178
+ res = self._http.request(method, url, headers=self._headers(), json=body)
179
+ except httpx.HTTPError as e:
180
+ err = e
181
+ out = self._outcome(res, err, url)
182
+ if not isinstance(out, Exception):
183
+ return out
184
+ last = out
185
+ if not isinstance(out, UnavailableError):
186
+ raise out
187
+ self._next += 1
188
+ raise last
189
+
190
+ def upsert(self, documents: Sequence[Document]) -> WriteResult:
191
+ """Inserts or replaces documents."""
192
+ body: dict[str, Any] = {"documents": list(documents)}
193
+ if self.token:
194
+ body["after"] = self.token
195
+ r = self._call("POST", f"{self._base}/documents", body)
196
+ self.observe(r["consistency_token"])
197
+ return WriteResult(count=r["count"], token=r["consistency_token"])
198
+
199
+ def patch(self, id: Id, set: Mapping[str, Any]) -> WriteResult:
200
+ """Changes some fields of one document (``None`` clears a field). A missing document
201
+ is not created. ``count`` is the number of documents changed."""
202
+ return self.patch_many([{"id": id, "set": dict(set)}])
203
+
204
+ def patch_many(self, patches: Sequence[Mapping[str, Any]]) -> WriteResult:
205
+ """Changes fields of several documents: ``[{"id": ..., "set": {...}}]``."""
206
+ body: dict[str, Any] = {"patches": [dict(p) for p in patches]}
207
+ if self.token:
208
+ body["after"] = self.token
209
+ r = self._call("POST", f"{self._base}/documents/patch", body)
210
+ self.observe(r["consistency_token"])
211
+ return WriteResult(count=r["patched"], token=r["consistency_token"])
212
+
213
+ def get(self, id: Id, *, consistency: Optional[Consistency] = None) -> Optional[Document]:
214
+ """One document, or ``None`` when there is none (or it was taken down)."""
215
+ try:
216
+ return self._call("GET", doc_path(id, self.token, consistency, self._base))
217
+ except NotFoundError:
218
+ return None
219
+
220
+ def search(
221
+ self,
222
+ *,
223
+ k: Optional[int] = None,
224
+ vector: Union[VectorLeg, Sequence[VectorLeg], None] = None,
225
+ text: Optional[str] = None,
226
+ text_field: Optional[str] = None,
227
+ all_terms: bool = False,
228
+ filter: Optional[Filter] = None,
229
+ fusion: Optional[Mapping[str, Any]] = None,
230
+ oversample: Optional[int] = None,
231
+ with_documents: Optional[bool] = None,
232
+ consistency: Optional[Consistency] = None,
233
+ group_by: Optional[str] = None,
234
+ ) -> list[Hit]:
235
+ """Hybrid search: vector legs, a text leg and a filter, fused. ``group_by`` returns one
236
+ hit per value of that field (each document once, at its best chunk)."""
237
+ body = search_body(
238
+ k=k, vector=vector, text=text, text_field=text_field, all_terms=all_terms, filter=filter,
239
+ fusion=fusion, oversample=oversample, with_documents=with_documents,
240
+ consistency=consistency, after=self.token, group_by=group_by,
241
+ )
242
+ return [Hit.from_json(h) for h in self._call("POST", f"{self._base}/search", body)["hits"]]
243
+
244
+ def delete(
245
+ self,
246
+ *,
247
+ ids: Optional[Sequence[Id]] = None,
248
+ filter: Optional[Filter] = None,
249
+ parent: Optional[Id] = None,
250
+ ) -> DeleteResult:
251
+ """Takes documents down: by ``ids``, by ``filter`` (every document that matches when
252
+ the deletion is applied), by ``parent`` (a document's chunks), or ``ids`` that also
253
+ match ``filter``."""
254
+ body = delete_body(ids=ids, filter=filter, parent=parent, parent_field=self._parent_field, after=self.token)
255
+ r = self._call("POST", f"{self._base}/documents/delete", body)
256
+ self.observe(r["consistency_token"])
257
+ return delete_result(r)
258
+
259
+ def forget_tenant(self, tenant: str) -> DeleteResult:
260
+ """Erases a tenant: every document it holds (unscoped keys with the takedown role)."""
261
+ q = f"?after={urllib.parse.quote(self.token)}" if self.token else ""
262
+ r = self._call("DELETE", f"{self._base}/tenants/{urllib.parse.quote(tenant, safe='')}{q}")
263
+ self.observe(r["consistency_token"])
264
+ return delete_result(r)
265
+
266
+ def schema(self) -> dict[str, Any]:
267
+ """The collection schema (reserved fields hidden)."""
268
+ return self._call("GET", f"{self._base}/schema")
269
+
270
+
271
+ class AsyncClient(_Base):
272
+ """The same API as :class:`Client`, with ``async`` methods."""
273
+
274
+ def __init__(self, url: Union[str, Sequence[str]], api_key: Optional[str] = None, **kw: Any) -> None:
275
+ super().__init__(url, api_key, **kw)
276
+ self._http = httpx.AsyncClient(timeout=self._timeout)
277
+
278
+ async def aclose(self) -> None:
279
+ await self._http.aclose()
280
+
281
+ async def __aenter__(self) -> "AsyncClient":
282
+ return self
283
+
284
+ async def __aexit__(self, *exc: object) -> None:
285
+ await self.aclose()
286
+
287
+ def with_tenant(self, tenant: str) -> "AsyncClient":
288
+ """A view acting for ``tenant`` (unscoped keys), sharing this client's token."""
289
+ return AsyncClient(**self._view_args(tenant=tenant))
290
+
291
+ def collection(self, name: str) -> "AsyncClient":
292
+ """A view acting on collection ``name``, sharing this client's token and tenant."""
293
+ return AsyncClient(**self._view_args(collection=name))
294
+
295
+ async def create_collection(
296
+ self,
297
+ name: str,
298
+ schema: Mapping[str, Any],
299
+ *,
300
+ shards: Optional[int] = None,
301
+ expires_field: Optional[str] = None,
302
+ ) -> dict[str, Any]:
303
+ body: dict[str, Any] = {"name": name, "schema": dict(schema)}
304
+ if shards is not None:
305
+ body["shards"] = shards
306
+ if expires_field:
307
+ body["expires_field"] = expires_field
308
+ return await self._call("POST", "/v1/collections", body)
309
+
310
+ async def list_collections(self) -> list[dict[str, Any]]:
311
+ return (await self._call("GET", "/v1/collections"))["collections"]
312
+
313
+ async def drop_collection(self, name: str) -> None:
314
+ await self._call("DELETE", "/v1/collections/" + urllib.parse.quote(name, safe=""))
315
+
316
+ async def _call(self, method: str, path: str, body: Any = None) -> Any:
317
+ last: Exception = UnavailableError("no attempt", 0)
318
+ for attempt in range(self._retries + 1):
319
+ if attempt:
320
+ await asyncio.sleep(0.1 * 2 ** (attempt - 1))
321
+ url = self._urls[self._next % len(self._urls)] + path
322
+ res, err = None, None
323
+ try:
324
+ res = await self._http.request(method, url, headers=self._headers(), json=body)
325
+ except httpx.HTTPError as e:
326
+ err = e
327
+ out = self._outcome(res, err, url)
328
+ if not isinstance(out, Exception):
329
+ return out
330
+ last = out
331
+ if not isinstance(out, UnavailableError):
332
+ raise out
333
+ self._next += 1
334
+ raise last
335
+
336
+ async def upsert(self, documents: Sequence[Document]) -> WriteResult:
337
+ body: dict[str, Any] = {"documents": list(documents)}
338
+ if self.token:
339
+ body["after"] = self.token
340
+ r = await self._call("POST", f"{self._base}/documents", body)
341
+ self.observe(r["consistency_token"])
342
+ return WriteResult(count=r["count"], token=r["consistency_token"])
343
+
344
+ async def patch(self, id: Id, set: Mapping[str, Any]) -> WriteResult:
345
+ return await self.patch_many([{"id": id, "set": dict(set)}])
346
+
347
+ async def patch_many(self, patches: Sequence[Mapping[str, Any]]) -> WriteResult:
348
+ body: dict[str, Any] = {"patches": [dict(p) for p in patches]}
349
+ if self.token:
350
+ body["after"] = self.token
351
+ r = await self._call("POST", f"{self._base}/documents/patch", body)
352
+ self.observe(r["consistency_token"])
353
+ return WriteResult(count=r["patched"], token=r["consistency_token"])
354
+
355
+ async def get(self, id: Id, *, consistency: Optional[Consistency] = None) -> Optional[Document]:
356
+ try:
357
+ return await self._call("GET", doc_path(id, self.token, consistency, self._base))
358
+ except NotFoundError:
359
+ return None
360
+
361
+ async def search(
362
+ self,
363
+ *,
364
+ k: Optional[int] = None,
365
+ vector: Union[VectorLeg, Sequence[VectorLeg], None] = None,
366
+ text: Optional[str] = None,
367
+ text_field: Optional[str] = None,
368
+ all_terms: bool = False,
369
+ filter: Optional[Filter] = None,
370
+ fusion: Optional[Mapping[str, Any]] = None,
371
+ oversample: Optional[int] = None,
372
+ with_documents: Optional[bool] = None,
373
+ consistency: Optional[Consistency] = None,
374
+ group_by: Optional[str] = None,
375
+ ) -> list[Hit]:
376
+ body = search_body(
377
+ k=k, vector=vector, text=text, text_field=text_field, all_terms=all_terms, filter=filter,
378
+ fusion=fusion, oversample=oversample, with_documents=with_documents,
379
+ consistency=consistency, after=self.token, group_by=group_by,
380
+ )
381
+ r = await self._call("POST", f"{self._base}/search", body)
382
+ return [Hit.from_json(h) for h in r["hits"]]
383
+
384
+ async def delete(
385
+ self,
386
+ *,
387
+ ids: Optional[Sequence[Id]] = None,
388
+ filter: Optional[Filter] = None,
389
+ parent: Optional[Id] = None,
390
+ ) -> DeleteResult:
391
+ body = delete_body(ids=ids, filter=filter, parent=parent, parent_field=self._parent_field, after=self.token)
392
+ r = await self._call("POST", f"{self._base}/documents/delete", body)
393
+ self.observe(r["consistency_token"])
394
+ return delete_result(r)
395
+
396
+ async def forget_tenant(self, tenant: str) -> DeleteResult:
397
+ q = f"?after={urllib.parse.quote(self.token)}" if self.token else ""
398
+ r = await self._call("DELETE", f"{self._base}/tenants/{urllib.parse.quote(tenant, safe='')}{q}")
399
+ self.observe(r["consistency_token"])
400
+ return delete_result(r)
401
+
402
+ async def schema(self) -> dict[str, Any]:
403
+ return await self._call("GET", f"{self._base}/schema")
File without changes
@@ -0,0 +1,130 @@
1
+ """Against a running node (clients/test-live.sh): CAIRN_URL, CAIRN_ADMIN_KEY, CAIRN_KEY (read,write,takedown),
2
+ CAIRN_ACME_KEY (scoped to tenant "acme"), CAIRN_READ_KEY (read only). Skipped without them."""
3
+
4
+ import asyncio
5
+ import os
6
+ import time
7
+
8
+ import pytest
9
+
10
+ import cairn_db as c
11
+
12
+ URL = os.environ.get("CAIRN_URL")
13
+ pytestmark = pytest.mark.skipif(not URL, reason="CAIRN_URL not set")
14
+ RUN = f"py-{os.getpid()}-{int(time.time() * 1000)}"
15
+
16
+
17
+ def vec(i):
18
+ return [1.0, float(i), 0.5, 0.0]
19
+
20
+
21
+ def test_client_against_a_live_node():
22
+ root = c.Client(URL.split(","), os.environ["CAIRN_KEY"])
23
+ db = root.with_tenant(RUN) # a fresh tenant: this run's data only
24
+ chunks = [{"id": f"doc-1#{i}", "parent": "doc-1", "text": f"ocelot chunk {i}", "embedding": vec(i), "n": i} for i in range(4)]
25
+ chunks += [
26
+ {"id": "doc-2#0", "parent": "doc-2", "text": "ocelot other", "embedding": vec(9), "n": 9, "tags": ["x"]},
27
+ {"id": 42, "text": "integer id", "embedding": vec(4), "n": 42},
28
+ {"id": "007", "text": "digits text id", "embedding": vec(5), "n": 7},
29
+ ]
30
+ assert db.upsert(chunks).count == 7
31
+
32
+ assert db.get("doc-1#2") == {"id": "doc-1#2", "parent": "doc-1", "text": "ocelot chunk 2", "embedding": vec(2), "n": 2}
33
+ assert db.get(42)["text"] == "integer id"
34
+ assert db.get("007")["text"] == "digits text id"
35
+ assert db.get("missing") is None
36
+ hits = db.search(k=20, text="ocelot", text_field="text")
37
+ assert sorted(h.id for h in hits) == ["doc-1#0", "doc-1#1", "doc-1#2", "doc-1#3", "doc-2#0"]
38
+ assert all(h.document and h.tenant is None for h in hits)
39
+ hits = db.search(k=3, vector={"field": "embedding", "values": vec(4)}, filter=c.range_("n", gte=4), with_documents=False)
40
+ assert hits[0].id == 42
41
+
42
+ other = c.Client(URL, os.environ["CAIRN_KEY"], tenant=RUN, token=db.token)
43
+ assert other.get("doc-2#0")["parent"] == "doc-2"
44
+
45
+ assert db.patch("doc-1#2", {"text": "ocelot patched", "n": None}).count == 1
46
+ assert db.get("doc-1#2") == {"id": "doc-1#2", "parent": "doc-1", "text": "ocelot patched", "embedding": vec(2)}
47
+ assert db.patch_many([{"id": 42, "set": {"n": 43}}, {"id": "nope", "set": {"n": 1}}]).count == 1
48
+ assert db.get(42)["n"] == 43
49
+ with pytest.raises(c.InvalidInputError):
50
+ db.patch(42, {"nope": 1})
51
+ assert db.delete(parent="doc-1").deleted == 4
52
+ assert db.get("doc-1#0") is None
53
+ assert [h.id for h in db.search(k=20, text="ocelot", text_field="text")] == ["doc-2#0"]
54
+ assert db.delete(ids=[42, "007"], filter=c.eq("n", 7)).deleted == 1
55
+ assert db.get("007") is None and db.get(42)["n"] == 43
56
+ assert db.delete(ids=[42]).count == 1
57
+ assert db.get(42) is None
58
+
59
+ with pytest.raises(c.InvalidInputError):
60
+ db.upsert([{"id": 1, "nope": 1}])
61
+ with pytest.raises(c.AuthenticationError):
62
+ c.Client(URL, "cairn_wrong_key").schema()
63
+ with pytest.raises(c.ForbiddenError):
64
+ c.Client(URL, os.environ["CAIRN_READ_KEY"]).upsert([{"id": 1}])
65
+
66
+ acme = c.Client(URL, os.environ["CAIRN_ACME_KEY"])
67
+ acme.upsert([{"id": "shared", "text": "acme's", "embedding": vec(1)}])
68
+ db.upsert([{"id": "shared", "text": "run's", "embedding": vec(1)}])
69
+ assert acme.get("shared")["text"] == "acme's"
70
+ assert db.get("shared")["text"] == "run's"
71
+ with pytest.raises(c.ForbiddenError):
72
+ acme.with_tenant("globex").get("shared")
73
+ with pytest.raises(c.ForbiddenError):
74
+ acme.forget_tenant("acme")
75
+
76
+ assert root.forget_tenant(RUN).deleted == 2
77
+ assert db.search(k=10) == []
78
+ assert root.forget_tenant("acme").deleted >= 1
79
+ assert acme.get("shared") is None
80
+
81
+
82
+ def test_async_client_against_a_live_node():
83
+ async def run():
84
+ async with c.AsyncClient(URL, os.environ["CAIRN_KEY"], tenant=RUN + "-async") as db:
85
+ await db.upsert([{"id": f"a{i}", "parent": "p", "text": "lynx", "embedding": vec(i)} for i in range(3)])
86
+ assert (await db.get("a1"))["text"] == "lynx"
87
+ assert len(await db.search(text="lynx", text_field="text")) == 3
88
+ assert (await db.delete(parent="p")).deleted == 3
89
+ assert await db.search(text="lynx", text_field="text") == []
90
+
91
+ asyncio.run(run())
92
+
93
+
94
+ def test_collections_against_a_live_node():
95
+ admin = c.Client(URL, os.environ["CAIRN_ADMIN_KEY"])
96
+ name = f"py-notes-{os.getpid()}"
97
+ schema = {"fields": [{"name": "embedding", "kind": {"Vector": {"dims": 2, "metric": "L2"}}},
98
+ {"name": "body", "kind": "Text"}, {"name": "parent", "kind": "Enum"}]}
99
+ assert admin.create_collection(name, schema, shards=2) == {"name": name, "shards": 2, "schema": schema}
100
+ assert name in [x["name"] for x in admin.list_collections()]
101
+ app = c.Client(URL, os.environ["CAIRN_KEY"])
102
+ notes = app.collection(name)
103
+ notes.upsert([{"id": f"n{i}", "embedding": [float(i), 1.0], "body": f"walrus {i}", "parent": f"p{i % 2}"} for i in range(4)])
104
+ assert notes.get("n1")["body"] == "walrus 1"
105
+ assert app.get("n1") is None, "not in default"
106
+ assert len(notes.search(text="walrus", text_field="body")) == 4
107
+ grouped = notes.search(vector={"field": "embedding", "values": [3.0, 1.0]}, group_by="parent")
108
+ assert [(h.id, h.group) for h in grouped] == [("n3", "p1"), ("n2", "p0")]
109
+ assert notes.delete(parent="p0").deleted == 2
110
+ assert sorted(h.id for h in notes.search(k=10)) == ["n1", "n3"]
111
+ assert notes.schema() == schema
112
+ with pytest.raises(c.ForbiddenError):
113
+ app.create_collection("x", schema)
114
+ admin.drop_collection(name)
115
+ assert name not in [x["name"] for x in admin.list_collections()]
116
+ with pytest.raises(c.NotFoundError):
117
+ notes.search(k=1)
118
+
119
+
120
+ def test_retention_through_the_client():
121
+ admin = c.Client(URL, os.environ["CAIRN_ADMIN_KEY"])
122
+ name = f"py-ttl-{os.getpid()}"
123
+ schema = {"fields": [{"name": "body", "kind": "Text"}, {"name": "until", "kind": "I64"}]}
124
+ assert admin.create_collection(name, schema, expires_field="until")["expires_field"] == "until"
125
+ col = c.Client(URL, os.environ["CAIRN_KEY"]).collection(name)
126
+ now = int(time.time() * 1000)
127
+ col.upsert([{"id": "gone", "body": "x", "until": now - 1000}, {"id": "kept", "body": "x", "until": now + 600_000}])
128
+ assert col.get("gone") is None
129
+ assert [h.id for h in col.search(k=10)] == ["kept"]
130
+ admin.drop_collection(name)
@@ -0,0 +1,109 @@
1
+ """Request shapes, token tracking, retries and typed errors, on a mock transport."""
2
+
3
+ import asyncio
4
+ import json
5
+
6
+ import httpx
7
+ import pytest
8
+
9
+ import cairn_db as c
10
+
11
+
12
+ def mock(db, responses, calls):
13
+ def handler(req: httpx.Request) -> httpx.Response:
14
+ calls.append(req)
15
+ r = responses.pop(0)
16
+ if isinstance(r, Exception):
17
+ raise r
18
+ status, body = r
19
+ return httpx.Response(status, json=body)
20
+
21
+ if isinstance(db, c.AsyncClient):
22
+ db._http = httpx.AsyncClient(transport=httpx.MockTransport(handler))
23
+ else:
24
+ db._http = httpx.Client(transport=httpx.MockTransport(handler))
25
+ return db
26
+
27
+
28
+ def test_tokens_and_filters():
29
+ assert c.merge_tokens("0.5,1.2", "1.7,2.1", None, "") == "0.5,1.7,2.1"
30
+ with pytest.raises(c.InvalidInputError):
31
+ c.merge_tokens("x")
32
+ assert c.and_(c.eq("s", "tv"), c.or_(c.in_("t", ["a"]), c.not_(c.is_null("d"))), c.range_("n", gte=1, lt=5)) == {
33
+ "and": [
34
+ {"field": "s", "eq": "tv"},
35
+ {"or": [{"field": "t", "in": ["a"]}, {"not": {"field": "d", "is_null": True}}]},
36
+ {"field": "n", "gte": 1, "lt": 5},
37
+ ]
38
+ }
39
+
40
+
41
+ def test_requests_and_token_tracking():
42
+ calls = []
43
+ db = mock(
44
+ c.Client("http://n1/", "k", parent_field="doc"),
45
+ [
46
+ (200, {"count": 2, "consistency_token": "0.4,1.9"}),
47
+ (200, {"id": "12"}),
48
+ (200, {"hits": [{"id": "a", "score": 0.5, "legs": [{"rank": 1, "score": 2.0}, None], "_tenant": "acme"}]}),
49
+ (200, {"deleted": 3, "consistency_token": "1.12"}),
50
+ (404, {"error": "no document"}),
51
+ (200, {"deleted": 5, "consistency_token": "0.20"}),
52
+ ],
53
+ calls,
54
+ )
55
+ assert db.upsert([{"id": "a"}, {"id": 3}]) == c.WriteResult(count=2, token="0.4,1.9")
56
+ assert calls[0].headers["authorization"] == "Bearer k"
57
+ db.get("12")
58
+ assert str(calls[1].url) == "http://n1/v1/documents/12?id_type=text&after=0.4%2C1.9"
59
+ hits = db.search(k=5, text="q", text_field="text", all_terms=True, vector=[{"field": "e", "values": [1.0]}],
60
+ with_documents=False, filter=c.eq("a", 1))
61
+ assert json.loads(calls[2].content) == {
62
+ "k": 5, "vectors": [{"field": "e", "values": [1.0]}], "text": {"field": "text", "query": "q", "all_terms": True},
63
+ "filter": {"field": "a", "eq": 1}, "with_documents": False, "after": "0.4,1.9"}
64
+ assert hits == [c.Hit(id="a", score=0.5, legs=[c.LegHit(1, 2.0), None], tenant="acme")]
65
+ assert db.delete(parent="doc-9") == c.DeleteResult(token="1.12", deleted=3)
66
+ assert json.loads(calls[3].content) == {"filter": {"field": "doc", "eq": "doc-9"}, "after": "0.4,1.9"}
67
+ assert db.token == "0.4,1.12"
68
+ assert db.get("a/b c") is None
69
+ assert str(calls[4].url).startswith("http://n1/v1/documents/a%2Fb%20c?after=")
70
+ acme = db.with_tenant("acme")
71
+ acme._http = db._http
72
+ acme.forget_tenant("globex")
73
+ assert calls[5].headers["cairn-tenant"] == "acme"
74
+ assert str(calls[5].url) == "http://n1/v1/tenants/globex?after=0.4%2C1.12"
75
+ assert db.token == "0.20,1.12", "tenant views share the token"
76
+ with pytest.raises(c.InvalidInputError):
77
+ db.delete()
78
+ with pytest.raises(c.InvalidInputError):
79
+ db.delete(parent="x", ids=[1])
80
+
81
+
82
+ def test_retries_and_errors():
83
+ calls = []
84
+ db = mock(c.Client(["http://a", "http://b"], retries=3), [
85
+ (503, {"error": "no leader"}), httpx.ConnectError("refused"), (200, {"count": 1, "consistency_token": "0.1"})
86
+ ], calls)
87
+ db.upsert([{"id": 1}])
88
+ assert [r.url.host for r in calls] == ["a", "b", "a"]
89
+ for status, cls in [(400, c.InvalidInputError), (401, c.AuthenticationError), (403, c.ForbiddenError), (500, c.CairnError)]:
90
+ db = mock(c.Client("http://a"), [(status, {"error": "boom"})], [])
91
+ with pytest.raises(cls) as e:
92
+ db.upsert([{"id": 1}])
93
+ assert (e.value.status, e.value.message) == (status, "boom")
94
+ db = mock(c.Client("http://a", retries=1), [(503, {}), (503, {})], [])
95
+ with pytest.raises(c.UnavailableError):
96
+ db.search()
97
+
98
+
99
+ def test_async_client():
100
+ async def run():
101
+ calls = []
102
+ db = mock(c.AsyncClient("http://a"), [
103
+ (503, {}), (200, {"count": 1, "consistency_token": "2.3"}), (200, {"hits": []}),
104
+ ], calls)
105
+ assert (await db.upsert([{"id": "x"}])).token == "2.3"
106
+ assert await db.search(filter=c.eq("a", 1)) == []
107
+ assert json.loads(calls[2].content)["after"] == "2.3"
108
+ await db.aclose()
109
+ asyncio.run(run())