vvz-agent-memory 1.7.46__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,12 @@
1
+ """vvz-agent-memory — the agent-memory CLI of the prompt-factory contract.
2
+
3
+ Agents of a `file_access: local` project write and read their dialogue messages
4
+ through this tool. It is a wrapper over `ailogger-client`: while the AI Logger is
5
+ reachable every call goes to it; while it is not, messages are kept in
6
+ `.agent-memory/<session_id>/` as `.msg` + `.emb` file pairs and are later moved
7
+ into the logger with their original identity and time by `agent-memory import`.
8
+
9
+ The package version equals the contract version it serves.
10
+ """
11
+
12
+ __version__ = "1.7.46"
@@ -0,0 +1,304 @@
1
+ """`agent-memory` — the command line.
2
+
3
+ Every command prints exactly one JSON object on stdout:
4
+ ``{"success": true, "store": ..., "data": ...}`` or ``{"success": false, "error": {"code", "message"}}``.
5
+
6
+ Exit codes: 0 success; 1 the logger refused or failed; 2 usage; 3 file store unavailable;
7
+ 4 logger unreachable (and no fallback); 5 not found; 6 configuration invalid;
8
+ 7 embedding service unavailable where a vector is required.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import argparse
14
+ import asyncio
15
+ import json
16
+ import sys
17
+ import uuid
18
+ from pathlib import Path
19
+ from typing import Any
20
+
21
+ from . import __version__, config as cfg
22
+ from .filestore import FileStore, FileStoreUnavailable, project_root
23
+ from .logger import LoggerFailure, LoggerUnreachable
24
+ from .memory import Memory, NotFound, Unavailable
25
+
26
+ KINDS = ("human", "model", "tool", "system")
27
+ RELATION_KINDS = ("reply", "influence", "tool_result")
28
+ EXIT = {"LOGGER_FAILURE": 1, "USAGE": 2, "FILE_STORE_UNAVAILABLE": 3, "LOGGER_UNREACHABLE": 4,
29
+ "NOT_FOUND": 5, "CONFIG_INVALID": 6, "EMBEDDING_UNAVAILABLE": 7}
30
+
31
+
32
+ def _uuid4(value: str) -> str:
33
+ try:
34
+ parsed = uuid.UUID(value)
35
+ except ValueError as exc:
36
+ raise argparse.ArgumentTypeError(f"not a UUID: {value!r}") from exc
37
+ if parsed.version != 4 or str(parsed) != value.lower():
38
+ raise argparse.ArgumentTypeError(f"not a canonical UUID4: {value!r}")
39
+ return str(parsed)
40
+
41
+
42
+ def _relation(value: str) -> dict[str, str]:
43
+ kind, _, target = value.partition(":")
44
+ if kind not in RELATION_KINDS:
45
+ raise argparse.ArgumentTypeError(f"relation kind must be one of {list(RELATION_KINDS)}: {value!r}")
46
+ return {"to_message_id": _uuid4(target), "kind": kind}
47
+
48
+
49
+ def _attributes(value: str) -> dict[str, Any]:
50
+ try:
51
+ parsed = json.loads(value)
52
+ except json.JSONDecodeError as exc:
53
+ raise argparse.ArgumentTypeError(f"--attributes is not JSON: {exc}") from exc
54
+ if not isinstance(parsed, dict):
55
+ raise argparse.ArgumentTypeError("--attributes must be a JSON object")
56
+ return parsed
57
+
58
+
59
+ def _non_negative(value: str) -> int:
60
+ number = int(value)
61
+ if number < 0:
62
+ raise argparse.ArgumentTypeError("must be >= 0")
63
+ return number
64
+
65
+
66
+ def _positive(value: str) -> int:
67
+ number = int(value)
68
+ if number <= 0:
69
+ raise argparse.ArgumentTypeError("must be > 0")
70
+ return number
71
+
72
+
73
+ def _batch(value: str) -> int:
74
+ number = _positive(value)
75
+ if number > 1000:
76
+ raise argparse.ArgumentTypeError("the logger imports at most 1000 messages per call")
77
+ return number
78
+
79
+
80
+ def _bool(value: str) -> bool:
81
+ if value.lower() in ("true", "yes", "1"):
82
+ return True
83
+ if value.lower() in ("false", "no", "0"):
84
+ return False
85
+ raise argparse.ArgumentTypeError("expected true or false")
86
+
87
+
88
+ def _global_options(parser: argparse.ArgumentParser, suppress: bool) -> None:
89
+ """The same options before and after the subcommand; the subcommand's copy only overrides when given."""
90
+
91
+ def default(value: Any) -> Any:
92
+ return argparse.SUPPRESS if suppress else value
93
+
94
+ parser.add_argument("--config", default=default(None),
95
+ help=f"configuration document (default ${cfg.ENV_VAR} or {cfg.DEFAULT_PATH})")
96
+ parser.add_argument("--root", default=default(None),
97
+ help="project root holding .agent-memory/ (default: the git top level of the cwd)")
98
+ parser.add_argument("--store", choices=("auto", "logger", "file"), default=default("auto"),
99
+ help="auto (default): logger, file store only while the logger is unreachable")
100
+
101
+
102
+ def build_parser() -> argparse.ArgumentParser:
103
+ parser = argparse.ArgumentParser(
104
+ prog="agent-memory",
105
+ description="Agent dialogue memory: AI Logger, with the .agent-memory/ file store while it is unreachable.",
106
+ )
107
+ parser.add_argument("--version", action="version", version=f"agent-memory {__version__}")
108
+ _global_options(parser, suppress=False)
109
+ common = argparse.ArgumentParser(add_help=False)
110
+ _global_options(common, suppress=True)
111
+
112
+ class _Sub(argparse._SubParsersAction): # every subcommand also takes the global options
113
+ def add_parser(self, name: str, **kwargs: Any) -> argparse.ArgumentParser:
114
+ kwargs.setdefault("parents", [common])
115
+ return super().add_parser(name, **kwargs)
116
+
117
+ parser.register("action", "parsers", _Sub)
118
+ sub = parser.add_subparsers(dest="command", required=True)
119
+
120
+ w = sub.add_parser("write", help="append one message (append_message)")
121
+ w.add_argument("--session-id", required=True, type=_uuid4)
122
+ text = w.add_mutually_exclusive_group(required=True)
123
+ text.add_argument("--text")
124
+ text.add_argument("--text-file", help="file holding the text; '-' reads stdin")
125
+ w.add_argument("--sender-kind", required=True, choices=KINDS)
126
+ w.add_argument("--sender-identity", required=True)
127
+ w.add_argument("--receiver-kind", required=True, choices=KINDS)
128
+ w.add_argument("--receiver-identity", required=True)
129
+ w.add_argument("--reply-to", type=_uuid4, action="append", default=[], metavar="MESSAGE_ID",
130
+ help="a reply relation to this message (repeatable)")
131
+ w.add_argument("--relation", type=_relation, action="append", default=[], metavar="KIND:MESSAGE_ID",
132
+ help=f"a relation, kind one of {list(RELATION_KINDS)} (repeatable)")
133
+ w.add_argument("--attributes", type=_attributes, help="free-form JSON object")
134
+ w.add_argument("--priority", type=_non_negative)
135
+ w.add_argument("--corrects", type=_uuid4, metavar="MESSAGE_ID", help="the message this one corrects")
136
+ w.add_argument("--message-id", type=_uuid4, help="identifier to accept the message under (a retry)")
137
+
138
+ r = sub.add_parser("read", help="read one referenced message")
139
+ r.add_argument("--session-id", required=True, type=_uuid4)
140
+ r.add_argument("--message-id", required=True, type=_uuid4)
141
+ r.add_argument("--sequence-number", type=_non_negative)
142
+
143
+ p = sub.add_parser("preview", help="messages of one sender, one character of text each")
144
+ p.add_argument("--session-id", required=True, type=_uuid4)
145
+ p.add_argument("--sender-identity", required=True)
146
+ p.add_argument("--page-size", type=_positive, default=100)
147
+
148
+ ws = sub.add_parser("working-set", help="recent and closest messages around an anchor")
149
+ ws.add_argument("--session-id", required=True, type=_uuid4)
150
+ anchor = ws.add_mutually_exclusive_group(required=True)
151
+ anchor.add_argument("--anchor-message-id", type=_uuid4)
152
+ anchor.add_argument("--anchor-text")
153
+ ws.add_argument("--recency-count", type=_non_negative, default=10)
154
+ ws.add_argument("--similarity-count", type=_non_negative, default=10)
155
+
156
+ rel = sub.add_parser("relations", help="declared relations of a session")
157
+ rel.add_argument("--session-id", required=True, type=_uuid4)
158
+
159
+ e = sub.add_parser("embed", help="compute the missing .emb files of the file store")
160
+ e.add_argument("--session-id", type=_uuid4, action="append", default=[])
161
+
162
+ i = sub.add_parser("import", help="move file-store messages into the logger with their ids and times")
163
+ i.add_argument("--session-id", type=_uuid4, action="append", default=[],
164
+ help="session to import (repeatable; default: every session in the store)")
165
+ i.add_argument("--batch-size", type=_batch, default=500)
166
+ i.add_argument("--keep", action="store_true",
167
+ help="leave imported files in place (default: move them to .agent-memory/.imported/)")
168
+
169
+ c = sub.add_parser("config", help="generate or validate the configuration document")
170
+ csub = c.add_subparsers(dest="config_command", required=True)
171
+ g = csub.add_parser("generate", help="write a configuration document; every setting has an argument")
172
+ g.add_argument("--output", required=True, help="path of the document to write")
173
+ g.add_argument("--force", action="store_true", help="overwrite an existing document")
174
+ for side, port in (("logger", 8008), ("embed", 8001)):
175
+ g.add_argument(f"--{side}-protocol", choices=cfg.PROTOCOLS, default="mtls")
176
+ g.add_argument(f"--{side}-host", required=True)
177
+ g.add_argument(f"--{side}-port", type=int, default=port)
178
+ g.add_argument(f"--{side}-timeout", type=float, default=cfg.DEFAULT_TIMEOUT)
179
+ g.add_argument(f"--{side}-cert")
180
+ g.add_argument(f"--{side}-key")
181
+ g.add_argument(f"--{side}-ca")
182
+ g.add_argument(f"--{side}-check-hostname", type=_bool)
183
+ g.add_argument(f"--{side}-token")
184
+ g.add_argument(f"--{side}-token-header")
185
+ g.add_argument("--embed-model", default="BAAI/bge-m3", help="must be the logger's embedding model")
186
+ g.add_argument("--embed-dimension", type=_positive, default=1024)
187
+ csub.add_parser("validate", help="validate the configuration document", parents=[common])
188
+ return parser
189
+
190
+
191
+ def _emit(payload: dict[str, Any]) -> None:
192
+ json.dump(payload, sys.stdout, ensure_ascii=False, default=str)
193
+ sys.stdout.write("\n")
194
+
195
+
196
+ def _fail(code: str, message: str, **extra: Any) -> int:
197
+ _emit({"success": False, "error": {"code": code, "message": message, **extra}})
198
+ return EXIT[code]
199
+
200
+
201
+ def _config_command(args: argparse.Namespace) -> int:
202
+ if args.config_command == "generate":
203
+ out = Path(args.output).expanduser().resolve()
204
+ if out.exists() and not args.force:
205
+ return _fail("USAGE", f"{out} exists; pass --force to overwrite")
206
+ document = cfg.generate(args)
207
+ errors = cfg.validate(document, out.parent)
208
+ if errors:
209
+ return _fail("CONFIG_INVALID", "the generated document is invalid; nothing written", errors=errors)
210
+ out.parent.mkdir(parents=True, exist_ok=True)
211
+ out.write_text(json.dumps(document, indent=2) + "\n", encoding="utf-8")
212
+ _emit({"success": True, "data": {"path": str(out)}})
213
+ return 0
214
+ try:
215
+ loaded = cfg.load(args.config)
216
+ except cfg.ConfigProblem as exc:
217
+ return _fail("CONFIG_INVALID", "configuration document is invalid", errors=exc.errors)
218
+ _emit({"success": True, "data": {"path": str(loaded.path), "errors": []}})
219
+ return 0
220
+
221
+
222
+ def _read_text(args: argparse.Namespace) -> str:
223
+ if args.text is not None:
224
+ return args.text
225
+ if args.text_file == "-":
226
+ return sys.stdin.read()
227
+ return Path(args.text_file).read_text(encoding="utf-8")
228
+
229
+
230
+ async def _run(args: argparse.Namespace, memory: Memory) -> dict[str, Any]:
231
+ if args.command == "write":
232
+ relations = [{"to_message_id": m, "kind": "reply"} for m in args.reply_to] + args.relation
233
+ text = _read_text(args)
234
+ if not text:
235
+ raise ValueError("the text is empty")
236
+ return await memory.write(
237
+ session_id=args.session_id, text=text, sender_kind=args.sender_kind,
238
+ sender_identity=args.sender_identity, receiver_kind=args.receiver_kind,
239
+ receiver_identity=args.receiver_identity, priority=args.priority, message_id=args.message_id,
240
+ relations=relations, corrects_message_id=args.corrects, attributes=args.attributes,
241
+ )
242
+ if args.command == "read":
243
+ return await memory.read(session_id=args.session_id, message_id=args.message_id,
244
+ sequence_number=args.sequence_number)
245
+ if args.command == "preview":
246
+ return await memory.preview(session_id=args.session_id, sender_identity=args.sender_identity,
247
+ page_size=args.page_size)
248
+ if args.command == "working-set":
249
+ return await memory.working_set(
250
+ session_id=args.session_id, anchor_message_id=args.anchor_message_id, anchor_text=args.anchor_text,
251
+ recency_count=args.recency_count, similarity_count=args.similarity_count)
252
+ if args.command == "relations":
253
+ return await memory.relations(session_id=args.session_id)
254
+ if args.command == "embed":
255
+ return await memory.embed_missing(session_ids=args.session_id)
256
+ if args.command == "import":
257
+ return await memory.import_sessions(session_ids=args.session_id, batch_size=args.batch_size,
258
+ keep=args.keep)
259
+ raise AssertionError(args.command)
260
+
261
+
262
+ def main(argv: list[str] | None = None) -> int:
263
+ args = build_parser().parse_args(argv)
264
+ if args.command == "config":
265
+ return _config_command(args)
266
+ try:
267
+ loaded = cfg.load(args.config)
268
+ except cfg.ConfigProblem as exc:
269
+ return _fail("CONFIG_INVALID", "configuration document is invalid", errors=exc.errors)
270
+ files: FileStore | None = None
271
+ problem: str | None = None
272
+ try:
273
+ files = FileStore.open(project_root(args.root))
274
+ except FileStoreUnavailable as exc:
275
+ problem = str(exc)
276
+ if args.store == "file" and files is None:
277
+ return _fail("FILE_STORE_UNAVAILABLE", problem or "file store unavailable")
278
+ if args.command in ("embed", "import") and files is None:
279
+ return _fail("FILE_STORE_UNAVAILABLE", problem or "file store unavailable")
280
+ try:
281
+ result = asyncio.run(_run(args, Memory(loaded, args.store, files, problem)))
282
+ except ValueError as exc:
283
+ return _fail("USAGE", str(exc))
284
+ except NotFound as exc:
285
+ return _fail("NOT_FOUND", str(exc))
286
+ except FileStoreUnavailable as exc:
287
+ return _fail("FILE_STORE_UNAVAILABLE", str(exc))
288
+ except LoggerUnreachable as exc:
289
+ return _fail("LOGGER_UNREACHABLE", str(exc))
290
+ except Unavailable as exc:
291
+ return _fail("LOGGER_UNREACHABLE", str(exc), sides=exc.reasons)
292
+ except LoggerFailure as exc:
293
+ return _fail("LOGGER_FAILURE", str(exc), logger_code=exc.code)
294
+ code = 0
295
+ if args.command == "import" and any(s["failed"] for s in result["data"]["sessions"]):
296
+ code = EXIT["LOGGER_FAILURE"]
297
+ if args.command == "embed" and result["data"]["failed"]:
298
+ code = EXIT["EMBEDDING_UNAVAILABLE"]
299
+ _emit({"success": code == 0, **result})
300
+ return code
301
+
302
+
303
+ if __name__ == "__main__":
304
+ sys.exit(main())
@@ -0,0 +1,219 @@
1
+ """The configuration document of the agent-memory CLI: location, reading, validation, generation.
2
+
3
+ The document is JSON with exactly three sections:
4
+
5
+ - ``ailogger_client`` — the published ai-logger client section, read by
6
+ ``ailogger_client.config`` itself (protocol, server.host/port, client.timeout,
7
+ client.ssl or ssl cert/key/ca, ssl.check_hostname, auth.token/token_header);
8
+ - ``embedding_client`` — the same sectioned shape for the embedding service;
9
+ - ``embedding_model`` — ``model`` and ``dimension``; they must be the logger's own,
10
+ otherwise vectors written while the logger is down are not comparable with its own.
11
+
12
+ Relative certificate paths are resolved against the document's directory.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import json
18
+ import os
19
+ from collections.abc import Mapping
20
+ from dataclasses import dataclass
21
+ from pathlib import Path
22
+ from typing import Any
23
+
24
+ ENV_VAR = "AGENT_MEMORY_CONFIG"
25
+ DEFAULT_PATH = Path.home() / ".config" / "vvz-agent-memory" / "config.json"
26
+ SECTIONS = ("ailogger_client", "embedding_client", "embedding_model")
27
+ PROTOCOLS = ("http", "https", "mtls")
28
+ DEFAULT_TIMEOUT = 10.0
29
+
30
+
31
+ class ConfigProblem(Exception):
32
+ """The configuration document cannot be used; ``errors`` lists every problem found."""
33
+
34
+ def __init__(self, errors: list[dict[str, str]]) -> None:
35
+ super().__init__("; ".join(f"{e['section']}.{e['key']}: {e['message']}" for e in errors))
36
+ self.errors = errors
37
+
38
+
39
+ @dataclass(frozen=True)
40
+ class Config:
41
+ path: Path
42
+ document: Mapping[str, Any]
43
+
44
+ @property
45
+ def base_dir(self) -> str:
46
+ return str(self.path.parent)
47
+
48
+ @property
49
+ def logger_section(self) -> Mapping[str, Any]:
50
+ return self.document["ailogger_client"]
51
+
52
+ @property
53
+ def embedding_section(self) -> Mapping[str, Any]:
54
+ return self.document["embedding_client"]
55
+
56
+ @property
57
+ def model(self) -> str:
58
+ return str(self.document["embedding_model"]["model"])
59
+
60
+ @property
61
+ def dimension(self) -> int:
62
+ return int(self.document["embedding_model"]["dimension"])
63
+
64
+
65
+ def resolve_path(explicit: str | None) -> Path:
66
+ """``--config`` wins, then ``$AGENT_MEMORY_CONFIG``, then ``~/.config/vvz-agent-memory/config.json``."""
67
+ if explicit:
68
+ return Path(explicit).expanduser().resolve()
69
+ env = os.environ.get(ENV_VAR)
70
+ if env:
71
+ return Path(env).expanduser().resolve()
72
+ return DEFAULT_PATH
73
+
74
+
75
+ def load(explicit: str | None) -> Config:
76
+ path = resolve_path(explicit)
77
+ if not path.is_file():
78
+ raise ConfigProblem([{"section": "-", "key": "-", "message": f"configuration document not found: {path}"}])
79
+ try:
80
+ document = json.loads(path.read_text(encoding="utf-8"))
81
+ except json.JSONDecodeError as exc:
82
+ raise ConfigProblem([{"section": "-", "key": "-", "message": f"{path}: invalid JSON: {exc}"}]) from exc
83
+ errors = validate(document, path.parent)
84
+ if errors:
85
+ raise ConfigProblem(errors)
86
+ return Config(path=path, document=document)
87
+
88
+
89
+ def _get(section: Any, dotted: str) -> Any:
90
+ node = section
91
+ for part in dotted.split("."):
92
+ if not isinstance(node, Mapping):
93
+ return None
94
+ node = node.get(part)
95
+ return node
96
+
97
+
98
+ def _err(section: str, key: str, message: str) -> dict[str, str]:
99
+ return {"section": section, "key": key, "message": message}
100
+
101
+
102
+ def _check_client_section(name: str, section: Any, base_dir: Path) -> list[dict[str, str]]:
103
+ errors: list[dict[str, str]] = []
104
+ if not isinstance(section, Mapping):
105
+ return [_err(name, "-", "must be an object")]
106
+ protocol = section.get("protocol", "http")
107
+ if protocol not in PROTOCOLS:
108
+ errors.append(_err(name, "protocol", f"must be one of {list(PROTOCOLS)}, got {protocol!r}"))
109
+ host = _get(section, "server.host")
110
+ if not isinstance(host, str) or not host:
111
+ errors.append(_err(name, "server.host", "required non-empty string"))
112
+ port = _get(section, "server.port")
113
+ if not isinstance(port, int) or isinstance(port, bool) or not 0 < port < 65536:
114
+ errors.append(_err(name, "server.port", "required integer 1..65535"))
115
+ timeout = _get(section, "client.timeout")
116
+ if timeout is not None and (not isinstance(timeout, (int, float)) or isinstance(timeout, bool) or timeout <= 0):
117
+ errors.append(_err(name, "client.timeout", "must be a positive number"))
118
+ material = {}
119
+ for key in ("cert", "key", "ca"):
120
+ value = _get(section, f"client.ssl.{key}") or _get(section, f"ssl.{key}")
121
+ if value is not None and not isinstance(value, str):
122
+ errors.append(_err(name, f"ssl.{key}", "must be a path string"))
123
+ continue
124
+ if value:
125
+ material[key] = value
126
+ if not (base_dir / value).is_file():
127
+ errors.append(_err(name, f"ssl.{key}", f"file not found: {base_dir / value}"))
128
+ if protocol == "mtls":
129
+ for key in ("cert", "key", "ca"):
130
+ if key not in material:
131
+ errors.append(_err(name, f"ssl.{key}", "required when protocol is mtls"))
132
+ elif protocol == "http" and material:
133
+ errors.append(_err(name, "ssl", "certificate material is set but protocol is http"))
134
+ if ("cert" in material) != ("key" in material):
135
+ errors.append(_err(name, "ssl.cert", "ssl.cert and ssl.key are set together or not at all"))
136
+ check_hostname = _get(section, "ssl.check_hostname")
137
+ if check_hostname is not None and not isinstance(check_hostname, bool):
138
+ errors.append(_err(name, "ssl.check_hostname", "must be a boolean"))
139
+ token = _get(section, "auth.token")
140
+ if token is not None and not isinstance(token, str):
141
+ errors.append(_err(name, "auth.token", "must be a string"))
142
+ header = _get(section, "auth.token_header")
143
+ if header is not None and not token:
144
+ errors.append(_err(name, "auth.token_header", "set without auth.token"))
145
+ return errors
146
+
147
+
148
+ def validate(document: Any, base_dir: Path) -> list[dict[str, str]]:
149
+ """Every problem of the document; empty when it is usable."""
150
+ if not isinstance(document, Mapping):
151
+ return [_err("-", "-", "the document must be a JSON object")]
152
+ errors: list[dict[str, str]] = []
153
+ for name in document:
154
+ if name not in SECTIONS:
155
+ errors.append(_err(name, "-", f"unknown section; allowed: {list(SECTIONS)}"))
156
+ for name in SECTIONS:
157
+ if name not in document:
158
+ errors.append(_err(name, "-", "required section is missing"))
159
+ if "ailogger_client" in document:
160
+ errors += _check_client_section("ailogger_client", document["ailogger_client"], base_dir)
161
+ if not errors:
162
+ from ailogger_client.config import ConfigError, load_config_kwargs
163
+
164
+ try:
165
+ load_config_kwargs(document["ailogger_client"], str(base_dir))
166
+ except ConfigError as exc:
167
+ errors.append(_err("ailogger_client", "-", f"refused by ailogger_client.config: {exc}"))
168
+ if "embedding_client" in document:
169
+ errors += _check_client_section("embedding_client", document["embedding_client"], base_dir)
170
+ model = document.get("embedding_model")
171
+ if model is not None:
172
+ if not isinstance(model, Mapping):
173
+ errors.append(_err("embedding_model", "-", "must be an object"))
174
+ else:
175
+ if not isinstance(model.get("model"), str) or not model.get("model"):
176
+ errors.append(_err("embedding_model", "model", "required non-empty string"))
177
+ dim = model.get("dimension")
178
+ if not isinstance(dim, int) or isinstance(dim, bool) or dim <= 0:
179
+ errors.append(_err("embedding_model", "dimension", "required positive integer"))
180
+ return errors
181
+
182
+
183
+ def _client_section(protocol: str, host: str, port: int, timeout: float, cert: str | None, key: str | None,
184
+ ca: str | None, check_hostname: bool | None, token: str | None,
185
+ token_header: str | None) -> dict[str, Any]:
186
+ section: dict[str, Any] = {
187
+ "protocol": protocol,
188
+ "server": {"host": host, "port": port},
189
+ "client": {"timeout": timeout},
190
+ }
191
+ ssl: dict[str, Any] = {}
192
+ for name, value in (("cert", cert), ("key", key), ("ca", ca)):
193
+ if value:
194
+ ssl[name] = value
195
+ if check_hostname is not None:
196
+ ssl["check_hostname"] = check_hostname
197
+ if ssl:
198
+ section["ssl"] = ssl
199
+ if token:
200
+ section["auth"] = {"token": token, "token_header": token_header or "X-API-Key"}
201
+ return section
202
+
203
+
204
+ def generate(args: Any) -> dict[str, Any]:
205
+ """Build a document from the ``config generate`` arguments; every setting has an argument."""
206
+ return {
207
+ "ailogger_client": _client_section(
208
+ args.logger_protocol, args.logger_host, args.logger_port, args.logger_timeout,
209
+ args.logger_cert, args.logger_key, args.logger_ca, args.logger_check_hostname,
210
+ args.logger_token, args.logger_token_header,
211
+ ),
212
+ "embedding_client": _client_section(
213
+ args.embed_protocol, args.embed_host, args.embed_port, args.embed_timeout,
214
+ args.embed_cert or args.logger_cert, args.embed_key or args.logger_key,
215
+ args.embed_ca or args.logger_ca, args.embed_check_hostname, args.embed_token,
216
+ args.embed_token_header,
217
+ ),
218
+ "embedding_model": {"model": args.embed_model, "dimension": args.embed_dimension},
219
+ }
@@ -0,0 +1,93 @@
1
+ """Vectors for the file store, through the published embed_client, with the logger's model.
2
+
3
+ The section is reshaped exactly as ai-logger's own binding does it
4
+ (``ailogger/embedding_client.py``), so the same service answers with the same model.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import math
10
+ from collections.abc import Mapping, Sequence
11
+ from pathlib import Path
12
+ from typing import Any
13
+
14
+ from .config import DEFAULT_TIMEOUT, Config, _get
15
+
16
+
17
+ class EmbeddingUnavailable(Exception):
18
+ """The embedding service did not give a vector; the message is kept without `.emb`."""
19
+
20
+
21
+ def _config_dict(section: Mapping[str, Any], base_dir: str) -> dict[str, Any]:
22
+ protocol = str(section.get("protocol", "http"))
23
+ host = _get(section, "server.host")
24
+ port = _get(section, "server.port")
25
+ timeout = float(_get(section, "client.timeout") or DEFAULT_TIMEOUT)
26
+
27
+ def path(key: str) -> str | None:
28
+ value = _get(section, f"client.ssl.{key}") or _get(section, f"ssl.{key}")
29
+ return str(Path(base_dir) / value) if value else None
30
+
31
+ scheme = "http" if protocol == "http" else "https"
32
+ result: dict[str, Any] = {
33
+ "protocol": protocol,
34
+ "server": {"host": host, "port": port, "base_url": f"{scheme}://{host}:{port}"},
35
+ "client": {"timeout": timeout},
36
+ }
37
+ cert, key, ca = path("cert"), path("key"), path("ca")
38
+ ssl: dict[str, Any] = {}
39
+ if cert or key or ca:
40
+ ssl.update(enabled=True, cert_file=cert, key_file=key, ca_cert_file=ca)
41
+ check_hostname = _get(section, "ssl.check_hostname")
42
+ if check_hostname is not None:
43
+ ssl["check_hostname"] = check_hostname
44
+ if ssl:
45
+ result["ssl"] = ssl
46
+ token = _get(section, "auth.token")
47
+ if token:
48
+ result["auth"] = {
49
+ "method": "api_key",
50
+ "api_keys": {"default": token},
51
+ "api_key_header": _get(section, "auth.token_header") or "X-API-Key",
52
+ }
53
+ return result
54
+
55
+
56
+ async def embed(config: Config, text: str) -> dict[str, Any]:
57
+ """The `.emb` body for ``text``: ``{model, dimension, embedding}``.
58
+
59
+ Raises EmbeddingUnavailable on any failure of the service or a vector of the wrong
60
+ model or dimension — a vector that does not match the logger's is worse than none.
61
+ """
62
+ from embed_client import EmbeddingClient
63
+ from embed_client.response_parsers import extract_embeddings
64
+
65
+ section = config.embedding_section
66
+ timeout = float(_get(section, "client.timeout") or DEFAULT_TIMEOUT)
67
+ try:
68
+ client = EmbeddingClient.from_config_dict(_config_dict(section, config.base_dir))
69
+ async with client:
70
+ data = await client.embed([text], model=config.model, wait=True, wait_timeout=int(timeout) or 60)
71
+ except Exception as exc: # the service, its transport and its client all count as "no vector"
72
+ raise EmbeddingUnavailable(f"{type(exc).__name__}: {exc}") from exc
73
+ if not isinstance(data, Mapping):
74
+ raise EmbeddingUnavailable(f"unexpected answer type {type(data).__name__}")
75
+ vectors = data.get("embeddings")
76
+ if vectors is None:
77
+ vectors = extract_embeddings({"result": {"success": True, "data": data}})
78
+ if not vectors or not isinstance(vectors[0], Sequence):
79
+ raise EmbeddingUnavailable("the answer carries no vector")
80
+ vector = [float(x) for x in vectors[0]]
81
+ model = data.get("model", config.model)
82
+ if model != config.model:
83
+ raise EmbeddingUnavailable(f"service answered with model {model!r}, configured {config.model!r}")
84
+ if len(vector) != config.dimension:
85
+ raise EmbeddingUnavailable(f"vector dimension {len(vector)}, configured {config.dimension}")
86
+ return {"model": model, "dimension": len(vector), "embedding": vector}
87
+
88
+
89
+ def cosine(a: Sequence[float], b: Sequence[float]) -> float:
90
+ dot = sum(x * y for x, y in zip(a, b, strict=True))
91
+ na = math.sqrt(sum(x * x for x in a))
92
+ nb = math.sqrt(sum(y * y for y in b))
93
+ return dot / (na * nb) if na and nb else 0.0