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.
- cairn_db_client-0.3.2/.gitignore +5 -0
- cairn_db_client-0.3.2/PKG-INFO +102 -0
- cairn_db_client-0.3.2/README.md +84 -0
- cairn_db_client-0.3.2/pyproject.toml +30 -0
- cairn_db_client-0.3.2/src/cairn_db/__init__.py +37 -0
- cairn_db_client-0.3.2/src/cairn_db/_core.py +283 -0
- cairn_db_client-0.3.2/src/cairn_db/client.py +403 -0
- cairn_db_client-0.3.2/src/cairn_db/py.typed +0 -0
- cairn_db_client-0.3.2/tests/test_live.py +130 -0
- cairn_db_client-0.3.2/tests/test_unit.py +109 -0
|
@@ -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())
|