pgbee 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pgbee/__init__.py +1 -0
- pgbee/cli.py +231 -0
- pgbee/db.py +114 -0
- pgbee/extension.py +87 -0
- pgbee/installer.py +84 -0
- pgbee/jobs.py +58 -0
- pgbee/providers.py +380 -0
- pgbee/py.typed +0 -0
- pgbee/schema.py +76 -0
- pgbee/settings.py +56 -0
- pgbee/sql/0001_init.sql +915 -0
- pgbee/sql/0002_cost_view_numeric.sql +15 -0
- pgbee/sql/0003_update_column_resets_backoff.sql +47 -0
- pgbee/sql/0004_decision_backend.sql +191 -0
- pgbee/sql/0005_budget.sql +215 -0
- pgbee/sql/0006_claim_decision_siblings.sql +110 -0
- pgbee/sql/0007_incremental_backfill.sql +298 -0
- pgbee/sql/0008_worker_role.sql +65 -0
- pgbee/sql/0009_lineage_retention.sql +87 -0
- pgbee/sql/0010_restore_safe_extension.sql +100 -0
- pgbee/sql/0011_result_version_at_claim.sql +202 -0
- pgbee/sql/0012_configure_comment.sql +4 -0
- pgbee/sql/0013_token_prices.sql +47 -0
- pgbee/sql/0014_complete_job_lock_order.sql +98 -0
- pgbee/worker.py +451 -0
- pgbee-0.1.0.dist-info/METADATA +36 -0
- pgbee-0.1.0.dist-info/RECORD +30 -0
- pgbee-0.1.0.dist-info/WHEEL +4 -0
- pgbee-0.1.0.dist-info/entry_points.txt +3 -0
- pgbee-0.1.0.dist-info/licenses/LICENSE +202 -0
pgbee/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""pgbee: installer and reference worker for pgbee derived columns."""
|
pgbee/cli.py
ADDED
|
@@ -0,0 +1,231 @@
|
|
|
1
|
+
"""`pgbee` command line: install the extension, run the worker, inspect derived columns."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import logging
|
|
7
|
+
import signal
|
|
8
|
+
import sys
|
|
9
|
+
import time
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import Annotated, Any
|
|
12
|
+
|
|
13
|
+
import psycopg
|
|
14
|
+
import structlog
|
|
15
|
+
import typer
|
|
16
|
+
from psycopg.rows import dict_row
|
|
17
|
+
|
|
18
|
+
from pgbee import extension
|
|
19
|
+
from pgbee.db import WORKER_BACKENDS, Contract
|
|
20
|
+
from pgbee.installer import InstalledAsExtension, install
|
|
21
|
+
from pgbee.providers import OpenAICompatibleProvider, OpenRouterProvider, Provider
|
|
22
|
+
from pgbee.settings import Settings, load_settings
|
|
23
|
+
from pgbee.worker import Worker, WorkerPool
|
|
24
|
+
|
|
25
|
+
app = typer.Typer(no_args_is_help=True, add_completion=False)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def _configure_logging(level: str) -> None:
|
|
29
|
+
numeric = logging.getLevelName(level.upper())
|
|
30
|
+
if not isinstance(numeric, int):
|
|
31
|
+
numeric = logging.INFO
|
|
32
|
+
structlog.configure(
|
|
33
|
+
wrapper_class=structlog.make_filtering_bound_logger(numeric),
|
|
34
|
+
processors=[
|
|
35
|
+
structlog.processors.TimeStamper(fmt="%H:%M:%S"),
|
|
36
|
+
structlog.processors.add_log_level,
|
|
37
|
+
structlog.dev.ConsoleRenderer(colors=sys.stderr.isatty()),
|
|
38
|
+
],
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
@app.command("install")
|
|
43
|
+
def install_cmd() -> None:
|
|
44
|
+
"""Apply the SQL files in sql/ that the database has not seen yet."""
|
|
45
|
+
settings = load_settings()
|
|
46
|
+
with psycopg.connect(settings.database_url) as conn:
|
|
47
|
+
try:
|
|
48
|
+
applied = install(conn)
|
|
49
|
+
except InstalledAsExtension as exc:
|
|
50
|
+
raise typer.BadParameter(str(exc)) from exc
|
|
51
|
+
if applied:
|
|
52
|
+
typer.echo(f"applied: {', '.join(str(v) for v in applied)}")
|
|
53
|
+
else:
|
|
54
|
+
typer.echo("up to date")
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
@app.command("extension-files")
|
|
58
|
+
def extension_files_cmd(
|
|
59
|
+
out_dir: Annotated[Path, typer.Argument(help="Where to write pgbee.control and the scripts.")],
|
|
60
|
+
) -> None:
|
|
61
|
+
"""Write the files for CREATE EXTENSION pgbee, to copy into `pg_config --sharedir`/extension."""
|
|
62
|
+
for path in extension.build(out_dir):
|
|
63
|
+
typer.echo(path)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
@app.command("status")
|
|
67
|
+
def status_cmd() -> None:
|
|
68
|
+
"""List derived columns with their version and queue counters."""
|
|
69
|
+
settings = load_settings()
|
|
70
|
+
with psycopg.connect(settings.database_url, row_factory=dict_row) as conn:
|
|
71
|
+
rows = conn.execute(
|
|
72
|
+
"SELECT c.table_schema, c.table_name, c.column_name, c.version, c.backend, c.model,"
|
|
73
|
+
" c.enabled, c.pending, c.claimed, c.done, c.dead, c.stale, c.human_overrides,"
|
|
74
|
+
" c.backfill_pending, c.backfill_scanned,"
|
|
75
|
+
" b.budget_usd, b.budget_period, b.spent_usd, b.exhausted"
|
|
76
|
+
" FROM bee.columns c JOIN bee.budgets b ON b.column_def_id = c.id"
|
|
77
|
+
" ORDER BY c.table_schema, c.table_name, c.column_name"
|
|
78
|
+
).fetchall()
|
|
79
|
+
if not rows:
|
|
80
|
+
typer.echo("no derived columns")
|
|
81
|
+
return
|
|
82
|
+
for r in rows:
|
|
83
|
+
state = "on" if r["enabled"] else "off"
|
|
84
|
+
typer.echo(
|
|
85
|
+
f"{r['table_schema']}.{r['table_name']}.{r['column_name']}"
|
|
86
|
+
f" v{r['version']} {r['backend']} {r['model']} [{state}]"
|
|
87
|
+
f" pending={r['pending']} claimed={r['claimed']} done={r['done']} dead={r['dead']}"
|
|
88
|
+
f" stale={r['stale']} human={r['human_overrides']}"
|
|
89
|
+
f"{_backfill(r)} {_spend(r)}"
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _backfill(r: dict[str, Any]) -> str:
|
|
94
|
+
if not r["backfill_pending"]:
|
|
95
|
+
return ""
|
|
96
|
+
return f" backfilling (scanned {r['backfill_scanned']} rows)"
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def _spend(r: dict[str, Any]) -> str:
|
|
100
|
+
spent = f"${r['spent_usd']:.4f}"
|
|
101
|
+
if r["budget_usd"] is None:
|
|
102
|
+
return f"spent {spent} this {r['budget_period']}, no cap"
|
|
103
|
+
flag = " EXHAUSTED" if r["exhausted"] else ""
|
|
104
|
+
return f"spent {spent} of ${r['budget_usd']} per {r['budget_period']}{flag}"
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
OPENAI_DEFAULT_URL = "https://api.openai.com/v1"
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def build_provider(settings: Settings) -> Provider:
|
|
111
|
+
try:
|
|
112
|
+
kind = settings.provider_kind()
|
|
113
|
+
except ValueError as exc:
|
|
114
|
+
raise typer.BadParameter(str(exc)) from exc
|
|
115
|
+
if kind == "openrouter":
|
|
116
|
+
assert settings.openrouter_api_key is not None
|
|
117
|
+
return OpenRouterProvider(
|
|
118
|
+
settings.openrouter_api_key,
|
|
119
|
+
settings.openrouter_base_url,
|
|
120
|
+
timeout=settings.request_timeout_seconds,
|
|
121
|
+
)
|
|
122
|
+
# Local servers (Ollama, vLLM) take no key, but the SDK wants a non empty one.
|
|
123
|
+
return OpenAICompatibleProvider(
|
|
124
|
+
settings.openai_api_key or "not-needed",
|
|
125
|
+
settings.openai_base_url or OPENAI_DEFAULT_URL,
|
|
126
|
+
timeout=settings.request_timeout_seconds,
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
@app.command("run")
|
|
131
|
+
def run_cmd(
|
|
132
|
+
once: bool = typer.Option(False, "--once", help="Process one batch and exit."),
|
|
133
|
+
drain: bool = typer.Option(
|
|
134
|
+
False,
|
|
135
|
+
"--drain",
|
|
136
|
+
help="Process until nothing is ready, then exit: for cron and serverless schedulers.",
|
|
137
|
+
),
|
|
138
|
+
max_seconds: float | None = typer.Option(
|
|
139
|
+
None,
|
|
140
|
+
"--max-seconds",
|
|
141
|
+
min=1,
|
|
142
|
+
help="With --drain, stop claiming after this many seconds (the batch in flight finishes).",
|
|
143
|
+
),
|
|
144
|
+
batch_size: int = typer.Option(20, "--batch-size", min=1),
|
|
145
|
+
backends: str = typer.Option(
|
|
146
|
+
"",
|
|
147
|
+
"--backends",
|
|
148
|
+
help="Comma separated backends this worker serves, one lane each. Default: all the"
|
|
149
|
+
f" provider supports ({','.join(WORKER_BACKENDS)} on OpenRouter).",
|
|
150
|
+
),
|
|
151
|
+
) -> None:
|
|
152
|
+
"""Consume the queue: call the models and write the results back."""
|
|
153
|
+
if once and drain:
|
|
154
|
+
raise typer.BadParameter("--once and --drain are alternatives")
|
|
155
|
+
if max_seconds is not None and not drain:
|
|
156
|
+
raise typer.BadParameter("--max-seconds works with --drain")
|
|
157
|
+
settings = load_settings()
|
|
158
|
+
provider = build_provider(settings)
|
|
159
|
+
_configure_logging(settings.log_level)
|
|
160
|
+
chosen = [b.strip() for b in backends.split(",") if b.strip()] or list(provider.backends)
|
|
161
|
+
if set(chosen) - set(WORKER_BACKENDS):
|
|
162
|
+
raise typer.BadParameter(f"--backends takes a subset of {','.join(WORKER_BACKENDS)}")
|
|
163
|
+
unsupported = set(chosen) - set(provider.backends)
|
|
164
|
+
if unsupported:
|
|
165
|
+
raise typer.BadParameter(
|
|
166
|
+
f"{','.join(sorted(unsupported))} not available with {settings.provider_kind()}:"
|
|
167
|
+
" the decision backend needs OpenRouter"
|
|
168
|
+
)
|
|
169
|
+
asyncio.run(
|
|
170
|
+
_run(
|
|
171
|
+
settings,
|
|
172
|
+
provider,
|
|
173
|
+
once=once,
|
|
174
|
+
drain=drain,
|
|
175
|
+
max_seconds=max_seconds,
|
|
176
|
+
batch_size=batch_size,
|
|
177
|
+
backends=chosen,
|
|
178
|
+
)
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
async def _run(
|
|
183
|
+
settings: Settings,
|
|
184
|
+
provider: Provider,
|
|
185
|
+
*,
|
|
186
|
+
once: bool,
|
|
187
|
+
drain: bool,
|
|
188
|
+
max_seconds: float | None,
|
|
189
|
+
batch_size: int,
|
|
190
|
+
backends: list[str],
|
|
191
|
+
) -> None:
|
|
192
|
+
structlog.get_logger("pgbee").info("provider", kind=settings.provider_kind(), backends=backends)
|
|
193
|
+
options: dict[str, Any] = {
|
|
194
|
+
"batch_size": batch_size,
|
|
195
|
+
"poll_interval": settings.poll_interval_seconds,
|
|
196
|
+
"claim_timeout_seconds": settings.claim_timeout_seconds,
|
|
197
|
+
"maintenance_interval": settings.maintenance_interval_seconds,
|
|
198
|
+
}
|
|
199
|
+
if once:
|
|
200
|
+
contract = await Contract.connect(settings.database_url)
|
|
201
|
+
try:
|
|
202
|
+
worker = Worker(
|
|
203
|
+
contract, provider, worker_id=settings.worker_id, backends=backends, **options
|
|
204
|
+
)
|
|
205
|
+
stats = await worker.run_once()
|
|
206
|
+
finally:
|
|
207
|
+
await contract.close()
|
|
208
|
+
typer.echo(
|
|
209
|
+
f"claimed={stats.claimed} failed={stats.failed} "
|
|
210
|
+
+ " ".join(f"{k}={v}" for k, v in stats.outcomes.items())
|
|
211
|
+
)
|
|
212
|
+
return
|
|
213
|
+
pool = WorkerPool(
|
|
214
|
+
lambda: Contract.connect(settings.database_url),
|
|
215
|
+
provider,
|
|
216
|
+
worker_id=settings.worker_id,
|
|
217
|
+
backends=backends,
|
|
218
|
+
**options,
|
|
219
|
+
)
|
|
220
|
+
loop = asyncio.get_running_loop()
|
|
221
|
+
for sig in (signal.SIGINT, signal.SIGTERM):
|
|
222
|
+
loop.add_signal_handler(sig, pool.stop)
|
|
223
|
+
if not drain:
|
|
224
|
+
await pool.run_forever()
|
|
225
|
+
return
|
|
226
|
+
deadline = None if max_seconds is None else time.monotonic() + max_seconds
|
|
227
|
+
for backend, stats in (await pool.drain(deadline)).items():
|
|
228
|
+
typer.echo(
|
|
229
|
+
f"{backend}: claimed={stats.claimed} failed={stats.failed} "
|
|
230
|
+
+ " ".join(f"{k}={v}" for k, v in stats.outcomes.items())
|
|
231
|
+
)
|
pgbee/db.py
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
1
|
+
"""The worker's side of the SQL contract: claim, complete, fail, reclaim, listen."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
from collections.abc import Sequence
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
import psycopg
|
|
10
|
+
from psycopg.rows import DictRow, dict_row
|
|
11
|
+
|
|
12
|
+
from pgbee.jobs import Job
|
|
13
|
+
|
|
14
|
+
WORKER_BACKENDS = ["llm", "decision", "embedding"]
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Contract:
|
|
18
|
+
def __init__(self, conn: psycopg.AsyncConnection[DictRow]) -> None:
|
|
19
|
+
self._conn = conn
|
|
20
|
+
|
|
21
|
+
@classmethod
|
|
22
|
+
async def connect(cls, database_url: str) -> Contract:
|
|
23
|
+
conn = await psycopg.AsyncConnection.connect(
|
|
24
|
+
database_url, autocommit=True, row_factory=dict_row
|
|
25
|
+
)
|
|
26
|
+
return cls(conn)
|
|
27
|
+
|
|
28
|
+
async def close(self) -> None:
|
|
29
|
+
await self._conn.close()
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
def broken(self) -> bool:
|
|
33
|
+
"""True when the connection is lost: the worker must stop, not fail jobs."""
|
|
34
|
+
return self._conn.broken or self._conn.closed
|
|
35
|
+
|
|
36
|
+
async def claim(
|
|
37
|
+
self, worker_id: str, batch_size: int, backends: Sequence[str] = WORKER_BACKENDS
|
|
38
|
+
) -> list[Job]:
|
|
39
|
+
cur = await self._conn.execute(
|
|
40
|
+
"SELECT * FROM bee.claim_jobs(%s, %s, %s::bee.backend[])"
|
|
41
|
+
" ORDER BY column_def_id, job_id",
|
|
42
|
+
(worker_id, batch_size, list(backends)),
|
|
43
|
+
)
|
|
44
|
+
return [Job.from_row(row) for row in await cur.fetchall()]
|
|
45
|
+
|
|
46
|
+
async def complete(
|
|
47
|
+
self,
|
|
48
|
+
job: Job,
|
|
49
|
+
value: Any,
|
|
50
|
+
*,
|
|
51
|
+
confidence: float | None,
|
|
52
|
+
model: str,
|
|
53
|
+
usage: dict[str, Any],
|
|
54
|
+
latency_ms: int,
|
|
55
|
+
details: dict[str, Any] | None = None,
|
|
56
|
+
) -> str:
|
|
57
|
+
cur = await self._conn.execute(
|
|
58
|
+
"SELECT bee.complete_job(%s::bigint, %s, %s::jsonb, %s::real, %s, %s::jsonb, %s,"
|
|
59
|
+
" %s::jsonb) AS outcome",
|
|
60
|
+
(
|
|
61
|
+
job.job_id,
|
|
62
|
+
job.source_hash,
|
|
63
|
+
json.dumps(value),
|
|
64
|
+
confidence,
|
|
65
|
+
model,
|
|
66
|
+
json.dumps(usage),
|
|
67
|
+
latency_ms,
|
|
68
|
+
None if details is None else json.dumps(details),
|
|
69
|
+
),
|
|
70
|
+
)
|
|
71
|
+
row = await cur.fetchone()
|
|
72
|
+
assert row is not None
|
|
73
|
+
return str(row["outcome"])
|
|
74
|
+
|
|
75
|
+
async def fail(self, job: Job, error: str, *, retryable: bool) -> str | None:
|
|
76
|
+
cur = await self._conn.execute(
|
|
77
|
+
"SELECT bee.fail_job(%s::bigint, %s, %s) AS status",
|
|
78
|
+
(job.job_id, error[:2000], retryable),
|
|
79
|
+
)
|
|
80
|
+
row = await cur.fetchone()
|
|
81
|
+
assert row is not None
|
|
82
|
+
return None if row["status"] is None else str(row["status"])
|
|
83
|
+
|
|
84
|
+
async def reclaim_stale(self, timeout_seconds: int) -> int:
|
|
85
|
+
cur = await self._conn.execute(
|
|
86
|
+
"SELECT bee.reclaim_stale(make_interval(secs => %s)) AS n", (timeout_seconds,)
|
|
87
|
+
)
|
|
88
|
+
row = await cur.fetchone()
|
|
89
|
+
assert row is not None
|
|
90
|
+
return int(row["n"])
|
|
91
|
+
|
|
92
|
+
async def prune(self, limit: int) -> tuple[int, int]:
|
|
93
|
+
"""One maintenance batch: (results, jobs) deleted, each at most `limit`."""
|
|
94
|
+
cur = await self._conn.execute("SELECT results, jobs FROM bee.prune(%s)", (limit,))
|
|
95
|
+
row = await cur.fetchone()
|
|
96
|
+
assert row is not None
|
|
97
|
+
return int(row["results"]), int(row["jobs"])
|
|
98
|
+
|
|
99
|
+
async def listen(self) -> None:
|
|
100
|
+
await self._conn.execute("LISTEN bee_jobs")
|
|
101
|
+
|
|
102
|
+
async def wait_for_notify(self, max_wait: float) -> bool:
|
|
103
|
+
"""Block until a NOTIFY on bee_jobs arrives or max_wait seconds pass. True on notify.
|
|
104
|
+
|
|
105
|
+
Notifications received while the worker was busy are queued by psycopg; they are all
|
|
106
|
+
drained here, since one claim serves them all.
|
|
107
|
+
"""
|
|
108
|
+
received = False
|
|
109
|
+
async for _ in self._conn.notifies(timeout=max_wait, stop_after=1):
|
|
110
|
+
received = True
|
|
111
|
+
if received:
|
|
112
|
+
async for _ in self._conn.notifies(timeout=0):
|
|
113
|
+
pass
|
|
114
|
+
return received
|
pgbee/extension.py
ADDED
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
"""Build the files for `CREATE EXTENSION pgbee` from the same versioned SQL files the installer
|
|
2
|
+
applies, so the two ways of installing cannot drift apart.
|
|
3
|
+
|
|
4
|
+
Version 0.N is file NNNN: `pgbee--0.1.sql` is file 0001 and every later file becomes an update
|
|
5
|
+
script `pgbee--0.(N-1)--0.N.sql`. Postgres chains them, both on a fresh CREATE EXTENSION and on
|
|
6
|
+
ALTER EXTENSION pgbee UPDATE.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
|
|
13
|
+
from pgbee.installer import SqlFile, sql_files
|
|
14
|
+
|
|
15
|
+
NAME = "pgbee"
|
|
16
|
+
COMMENT = "Derived columns computed by models: queue, versions, lineage, human overrides, budgets"
|
|
17
|
+
|
|
18
|
+
GUARD = (
|
|
19
|
+
'\\echo Use "CREATE EXTENSION pgbee" or "ALTER EXTENSION pgbee UPDATE" to load this file.'
|
|
20
|
+
" \\quit\n"
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
# Tables and sequences of the extension hold user data (definitions, jobs, lineage, spend):
|
|
24
|
+
# pg_dump must dump their rows. schema_version is left out because the scripts fill it.
|
|
25
|
+
CONFIG_DUMP = """
|
|
26
|
+
DO $pgbee$
|
|
27
|
+
DECLARE
|
|
28
|
+
r record;
|
|
29
|
+
BEGIN
|
|
30
|
+
FOR r IN
|
|
31
|
+
SELECT c.oid::regclass AS rel
|
|
32
|
+
FROM pg_catalog.pg_class c
|
|
33
|
+
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace
|
|
34
|
+
JOIN pg_catalog.pg_depend d ON d.classid = 'pg_catalog.pg_class'::regclass AND d.objid = c.oid
|
|
35
|
+
AND d.deptype = 'e'
|
|
36
|
+
JOIN pg_catalog.pg_extension e ON e.oid = d.refobjid AND e.extname = 'pgbee'
|
|
37
|
+
WHERE n.nspname = 'bee' AND c.relkind IN ('r', 'S')
|
|
38
|
+
AND c.relname <> 'schema_version'
|
|
39
|
+
AND NOT (c.oid = ANY (coalesce(e.extconfig, '{}')))
|
|
40
|
+
LOOP
|
|
41
|
+
PERFORM pg_catalog.pg_extension_config_dump(r.rel, '');
|
|
42
|
+
END LOOP;
|
|
43
|
+
END $pgbee$;
|
|
44
|
+
"""
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def version_of(sql_file: SqlFile) -> str:
|
|
48
|
+
return f"0.{sql_file.version}"
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def script(sql_file: SqlFile) -> str:
|
|
52
|
+
body = sql_file.path.read_text(encoding="utf-8")
|
|
53
|
+
return (
|
|
54
|
+
f"{GUARD}\n{body.rstrip()}\n{CONFIG_DUMP}\n"
|
|
55
|
+
f"INSERT INTO bee.schema_version (version) VALUES ({sql_file.version});\n"
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def control(default_version: str) -> str:
|
|
60
|
+
return (
|
|
61
|
+
f"comment = '{COMMENT}'\n"
|
|
62
|
+
f"default_version = '{default_version}'\n"
|
|
63
|
+
"relocatable = false\n"
|
|
64
|
+
"superuser = true\n"
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def build(out_dir: Path, directory: Path | None = None) -> list[Path]:
|
|
69
|
+
"""Write the control file and one script per SQL file into out_dir. Returns the paths."""
|
|
70
|
+
files = sql_files(directory)
|
|
71
|
+
out_dir.mkdir(parents=True, exist_ok=True)
|
|
72
|
+
written: list[Path] = []
|
|
73
|
+
previous: SqlFile | None = None
|
|
74
|
+
for sql_file in files:
|
|
75
|
+
if previous is None:
|
|
76
|
+
name = f"{NAME}--{version_of(sql_file)}.sql"
|
|
77
|
+
else:
|
|
78
|
+
name = f"{NAME}--{version_of(previous)}--{version_of(sql_file)}.sql"
|
|
79
|
+
path = out_dir / name
|
|
80
|
+
path.write_text(script(sql_file), encoding="utf-8")
|
|
81
|
+
written.append(path)
|
|
82
|
+
previous = sql_file
|
|
83
|
+
assert previous is not None
|
|
84
|
+
control_path = out_dir / f"{NAME}.control"
|
|
85
|
+
control_path.write_text(control(version_of(previous)), encoding="utf-8")
|
|
86
|
+
written.append(control_path)
|
|
87
|
+
return written
|
pgbee/installer.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""Apply the versioned SQL files in `sql/` to a database and track them in bee.schema_version."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import os
|
|
6
|
+
import re
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
import psycopg
|
|
12
|
+
from psycopg.rows import tuple_row
|
|
13
|
+
|
|
14
|
+
SQL_FILE_RE = re.compile(r"^(\d{4})_[a-z0-9_]+\.sql$")
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class InstalledAsExtension(RuntimeError):
|
|
18
|
+
"""The database has pgbee as a Postgres extension; the file installer must not touch it."""
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True)
|
|
22
|
+
class SqlFile:
|
|
23
|
+
version: int
|
|
24
|
+
path: Path
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def sql_dir() -> Path:
|
|
28
|
+
"""Directory holding the extension SQL files, shipped inside the package.
|
|
29
|
+
|
|
30
|
+
Overridable with PGBEE_SQL_DIR. The repository's top level `sql` links here.
|
|
31
|
+
"""
|
|
32
|
+
env = os.environ.get("PGBEE_SQL_DIR")
|
|
33
|
+
if env:
|
|
34
|
+
return Path(env)
|
|
35
|
+
return Path(__file__).resolve().parent / "sql"
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def sql_files(directory: Path | None = None) -> list[SqlFile]:
|
|
39
|
+
directory = directory or sql_dir()
|
|
40
|
+
files: list[SqlFile] = []
|
|
41
|
+
for path in sorted(directory.glob("*.sql")):
|
|
42
|
+
match = SQL_FILE_RE.match(path.name)
|
|
43
|
+
if not match:
|
|
44
|
+
raise ValueError(f"unexpected SQL file name {path.name}: want NNNN_name.sql")
|
|
45
|
+
files.append(SqlFile(version=int(match.group(1)), path=path))
|
|
46
|
+
if not files:
|
|
47
|
+
raise FileNotFoundError(f"no SQL files in {directory}")
|
|
48
|
+
return files
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def applied_versions(conn: psycopg.Connection[Any]) -> set[int]:
|
|
52
|
+
exists = conn.execute(
|
|
53
|
+
"SELECT 1 FROM pg_tables WHERE schemaname = 'bee' AND tablename = 'schema_version'"
|
|
54
|
+
).fetchone()
|
|
55
|
+
if not exists:
|
|
56
|
+
return set()
|
|
57
|
+
cur = conn.cursor(row_factory=tuple_row)
|
|
58
|
+
rows = cur.execute("SELECT version FROM bee.schema_version").fetchall()
|
|
59
|
+
return {int(row[0]) for row in rows}
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def install(conn: psycopg.Connection[Any], directory: Path | None = None) -> list[int]:
|
|
63
|
+
"""Apply every SQL file not yet recorded, one transaction each. Returns applied versions.
|
|
64
|
+
|
|
65
|
+
Refuses a database where pgbee was installed as an extension: there the files arrive through
|
|
66
|
+
ALTER EXTENSION pgbee UPDATE, and applying them here would detach objects from it.
|
|
67
|
+
"""
|
|
68
|
+
if conn.execute("SELECT 1 FROM pg_extension WHERE extname = 'pgbee'").fetchone():
|
|
69
|
+
raise InstalledAsExtension(
|
|
70
|
+
"pgbee is installed as the extension pgbee:"
|
|
71
|
+
" upgrade it with ALTER EXTENSION pgbee UPDATE"
|
|
72
|
+
)
|
|
73
|
+
applied = applied_versions(conn)
|
|
74
|
+
done: list[int] = []
|
|
75
|
+
for sql_file in sql_files(directory):
|
|
76
|
+
if sql_file.version in applied:
|
|
77
|
+
continue
|
|
78
|
+
with conn.transaction():
|
|
79
|
+
conn.execute(sql_file.path.read_text(encoding="utf-8"))
|
|
80
|
+
conn.execute(
|
|
81
|
+
"INSERT INTO bee.schema_version (version) VALUES (%s)", (sql_file.version,)
|
|
82
|
+
)
|
|
83
|
+
done.append(sql_file.version)
|
|
84
|
+
return done
|
pgbee/jobs.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
"""Job payloads as returned by bee.claim_jobs, and the text the models see."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@dataclass(frozen=True)
|
|
11
|
+
class Job:
|
|
12
|
+
job_id: int
|
|
13
|
+
column_def_id: int
|
|
14
|
+
column_version_id: int
|
|
15
|
+
backend: str
|
|
16
|
+
model: str
|
|
17
|
+
prompt: str | None
|
|
18
|
+
output_type: str
|
|
19
|
+
output_schema: Any
|
|
20
|
+
backend_config: dict[str, Any]
|
|
21
|
+
config: dict[str, Any]
|
|
22
|
+
row_pk: dict[str, Any]
|
|
23
|
+
source_hash: bytes
|
|
24
|
+
source: dict[str, Any]
|
|
25
|
+
attempts: int
|
|
26
|
+
|
|
27
|
+
@classmethod
|
|
28
|
+
def from_row(cls, row: dict[str, Any]) -> Job:
|
|
29
|
+
return cls(
|
|
30
|
+
job_id=int(row["job_id"]),
|
|
31
|
+
column_def_id=int(row["column_def_id"]),
|
|
32
|
+
column_version_id=int(row["column_version_id"]),
|
|
33
|
+
backend=str(row["backend"]),
|
|
34
|
+
model=str(row["model"]),
|
|
35
|
+
prompt=row["prompt"],
|
|
36
|
+
output_type=str(row["output_type"]),
|
|
37
|
+
output_schema=row["output_schema"],
|
|
38
|
+
backend_config=dict(row["backend_config"] or {}),
|
|
39
|
+
config=dict(row["config"] or {}),
|
|
40
|
+
row_pk=dict(row["row_pk"]),
|
|
41
|
+
source_hash=bytes(row["source_hash"]),
|
|
42
|
+
source=dict(row["source"]),
|
|
43
|
+
attempts=int(row["attempts"]),
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
def source_text(self) -> str:
|
|
47
|
+
"""Render the sources for a model: raw value for one column, labeled block otherwise."""
|
|
48
|
+
if len(self.source) == 1:
|
|
49
|
+
return _as_text(next(iter(self.source.values())))
|
|
50
|
+
return "\n".join(f"{name}: {_as_text(value)}" for name, value in self.source.items())
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _as_text(value: Any) -> str:
|
|
54
|
+
if value is None:
|
|
55
|
+
return ""
|
|
56
|
+
if isinstance(value, str):
|
|
57
|
+
return value
|
|
58
|
+
return json.dumps(value, ensure_ascii=False)
|