sqljev 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.
sqljev/finetune.py ADDED
@@ -0,0 +1,367 @@
1
+ """Fine-tune Laya on labelled SQL rows, on one GPU (a free Colab T4 is enough), and publish the checkpoint.
2
+
3
+ sqljev dataset "$DB_URL" "SELECT subject, body, team FROM tickets" --label team \\
4
+ --choice "which team should handle this?" --test-fraction 0.2 -o tickets.jsonl
5
+ sqljev finetune tickets.train.jsonl --out checkpoints/tickets --epochs 3
6
+ sqljev eval tickets.test.jsonl --model checkpoints/tickets
7
+ sqljev publish checkpoints/tickets --repo your-org/laya-tickets # Hugging Face Hub
8
+
9
+ Training data is what `sqljev dataset` writes: one row per line with the state (the row, as the model sees it at
10
+ query time), the question built by core.laya_question and the expected answer, so the checkpoint learns exactly
11
+ the questions sqljev will ask it.
12
+
13
+ The training loop follows Laya's own fine-tuning notebook (Convai Innovations, Apache 2.0,
14
+ https://github.com/NandhaKishorM/laya/tree/main/notebooks): policy-gradient updates against a proper scoring
15
+ rule plus soft cross-entropy, then per-question-type temperature calibration on a held-out slice. It is
16
+ single-GPU here, with gradient checkpointing and mixed precision, so it fits a 16 GB T4.
17
+ """
18
+ import json
19
+ import os
20
+ import random
21
+ import re
22
+ import time
23
+
24
+ from .core import JevError, check_question, laya_question, to_row_json
25
+
26
+ LAYA_BASE = "convaiinnovations/laya"
27
+ SUBFOLDERS = {"english": None, "multilingual": "multilingual", "typed-decisions": "typed-decisions"}
28
+
29
+
30
+ # ---------------------------------------------------------------- examples
31
+
32
+ def examples_from_rows(rows, question, kind, options=None, label=None, columns=None, drop_nulls=True):
33
+ """Labelled rows (dicts) -> training/eval records. The label column never reaches the state. For choice
34
+ questions without options, the options are the distinct labels."""
35
+ rows = [r for r in rows if r.get(label) is not None]
36
+ if kind == "choice" and not options:
37
+ options = sorted({str(r[label]) for r in rows})
38
+ kind, opts = check_question(kind, options)
39
+ q = {"q": laya_question(kind, question, opts)}
40
+ out = []
41
+ for r in rows:
42
+ lab = r[label]
43
+ if kind == "noul":
44
+ expected = _truthy(lab)
45
+ elif kind == "choice":
46
+ expected = str(lab)
47
+ if expected not in opts:
48
+ continue
49
+ else:
50
+ expected = opts.index(str(lab)) if str(lab) in opts else int(lab)
51
+ view = {k: v for k, v in r.items() if k != label and (not columns or k in columns)}
52
+ out.append({"state": json.loads(to_row_json(view, drop_nulls)), "questions": q, "expected": {"q": expected}})
53
+ return out
54
+
55
+
56
+ def split(records, test_fraction=0.2, seed=0):
57
+ """Deterministic train/test split."""
58
+ idx = list(range(len(records)))
59
+ random.Random(seed).shuffle(idx)
60
+ n_test = int(len(records) * test_fraction)
61
+ test = set(idx[:n_test])
62
+ return [r for i, r in enumerate(records) if i not in test], [r for i, r in enumerate(records) if i in test]
63
+
64
+
65
+ def write_jsonl(records, path):
66
+ with open(path, "w") as f:
67
+ for r in records:
68
+ f.write(json.dumps(r, ensure_ascii=False, default=str) + "\n")
69
+
70
+
71
+ def read_jsonl(path):
72
+ with open(path) as f:
73
+ return [json.loads(line) for line in f if line.strip() and not line.startswith("#")]
74
+
75
+
76
+ def _truthy(v):
77
+ if isinstance(v, bool):
78
+ return v
79
+ s = str(v).strip().lower()
80
+ if s in ("1", "true", "t", "yes", "y"):
81
+ return True
82
+ if s in ("0", "false", "f", "no", "n"):
83
+ return False
84
+ raise JevError("sqljev: label %r is not a boolean" % (v,))
85
+
86
+
87
+ def accuracy(records, **settings):
88
+ """Share of expected answers the model gets right on `records` (a list or a JSONL path), with any engine
89
+ settings (model=<checkpoint>, device=..., backend=...). Returns {"accuracy", "decisions", "seconds"}."""
90
+ from .core import Jev
91
+ records = read_jsonl(records) if isinstance(records, str) else records
92
+ jev = Jev(**settings)
93
+ groups = {}
94
+ for rec in records:
95
+ for qid, q in rec["questions"].items():
96
+ crit = q.get("criteria")
97
+ opts = list(crit) if crit else None
98
+ groups.setdefault((q["type"], q["instructions"], json.dumps(opts)), []).append(
99
+ (rec["state"], rec["expected"][qid]))
100
+ t0, right, total = time.time(), 0, 0
101
+ for (kind, instr, opts_json), items in groups.items():
102
+ prefix = "Is it true that "
103
+ query = instr[len(prefix):-1] if kind == "noul" and instr.startswith(prefix) and instr.endswith("?") else instr
104
+ answers = jev.evaluate([s for s, _ in items], query, kind, json.loads(opts_json))
105
+ for (_, exp), a in zip(items, answers):
106
+ if kind == "noul":
107
+ right += (a["noul"] >= 0.5) == bool(exp)
108
+ elif kind == "choice":
109
+ right += a["choice"] == exp
110
+ else:
111
+ right += round(a["score"]) == int(exp)
112
+ total += 1
113
+ return {"accuracy": round(right / total, 4) if total else 0.0, "decisions": total,
114
+ "seconds": round(time.time() - t0, 1)}
115
+
116
+
117
+ # ---------------------------------------------------------------- training
118
+
119
+ def base_dir(base=LAYA_BASE):
120
+ """Local directory of a Laya checkpoint: a path, a named base checkpoint, or a Hub id."""
121
+ if os.path.isdir(base):
122
+ return base
123
+ from huggingface_hub import snapshot_download
124
+ sub = SUBFOLDERS.get(base, None)
125
+ repo = LAYA_BASE if base in SUBFOLDERS else base
126
+ d = snapshot_download(repo)
127
+ try:
128
+ from laya.agent import _fix_tokenizer_config
129
+ _fix_tokenizer_config(d)
130
+ except Exception: # noqa: BLE001 -- older/newer laya: the files are already usable
131
+ pass
132
+ return os.path.join(d, sub) if sub else d
133
+
134
+
135
+ def _items(records, tok, cfg):
136
+ from laya.common import QTYPES, build_sequence, render_options
137
+ items, skipped = [], 0
138
+ for rec in records:
139
+ for qid, q in rec["questions"].items():
140
+ t, crit = q["type"], q.get("criteria") or {}
141
+ exp = rec["expected"][qid]
142
+ if t == "noul":
143
+ target = [0.0, 1.0] if exp else [1.0, 0.0]
144
+ elif t == "choice":
145
+ keys = list(crit.keys()) if isinstance(crit, dict) else list(crit)
146
+ target = [float(k == exp) for k in keys]
147
+ else:
148
+ target = [float(i == int(exp)) for i in range(len(crit))]
149
+ seq, markers = build_sequence(tok, rec["state"], {"t": t, "ins": q["instructions"], "crit": crit},
150
+ cfg["max_len"], cfg["head_max_len"])
151
+ if len(markers) != len(render_options({"t": t, "crit": crit})) or sum(target) != 1:
152
+ skipped += 1
153
+ continue
154
+ items.append({"ids": seq, "markers": markers, "qtype": QTYPES[t], "target": target,
155
+ "label": target.index(1.0)})
156
+ return items, skipped
157
+
158
+
159
+ def _collate(items, pad_id):
160
+ import torch
161
+ n, L = len(items), max(len(it["ids"]) for it in items)
162
+ kmax = max(len(it["markers"]) for it in items)
163
+ ids = torch.full((n, L), pad_id, dtype=torch.long)
164
+ att = torch.zeros((n, L), dtype=torch.long)
165
+ mpos = torch.zeros((n, kmax), dtype=torch.long)
166
+ mmask = torch.zeros((n, kmax), dtype=torch.bool)
167
+ target = torch.zeros((n, kmax), dtype=torch.float32)
168
+ for i, it in enumerate(items):
169
+ ids[i, :len(it["ids"])] = torch.tensor(it["ids"])
170
+ att[i, :len(it["ids"])] = 1
171
+ k = len(it["markers"])
172
+ mpos[i, :k] = torch.tensor(it["markers"])
173
+ mmask[i, :k] = True
174
+ target[i, :k] = torch.tensor(it["target"])
175
+ return ids, att, mpos, mmask, target, torch.tensor([it["qtype"] for it in items])
176
+
177
+
178
+ def _fit_temperature(pairs):
179
+ import torch
180
+ if len(pairs) < 10:
181
+ return 1.0
182
+ kmax = max(len(z) for z, _ in pairs)
183
+ Z = torch.full((len(pairs), kmax), -1e4)
184
+ T = torch.zeros((len(pairs), kmax))
185
+ for i, (z, t) in enumerate(pairs):
186
+ Z[i, :len(z)] = torch.tensor(z)
187
+ T[i, :len(t)] = torch.tensor(t)
188
+ log_t = torch.zeros(1, requires_grad=True)
189
+ opt = torch.optim.LBFGS([log_t], lr=0.1, max_iter=100)
190
+
191
+ def closure():
192
+ opt.zero_grad()
193
+ loss = -(T * torch.log_softmax(Z / log_t.exp(), -1)).sum(-1).mean()
194
+ loss.backward()
195
+ return loss
196
+ opt.step(closure)
197
+ return float(torch.clamp(log_t.exp(), 0.5, 5.0).item())
198
+
199
+
200
+ def finetune(train, out_dir, base=LAYA_BASE, epochs=3, micro_batch=8, grad_accum=8, lr_encoder=2.5e-5,
201
+ lr_head=1e-4, group_size=4, train_layers=None, device=None, seed=0, log=print):
202
+ """Fine-tune a Laya checkpoint on `train` (records or a JSONL path) and save it to `out_dir` in Laya's
203
+ format, loadable with laya.load(out_dir) and usable as SQLJEV_MODEL=out_dir.
204
+
205
+ train_layers: train only the top N encoder layers plus the decision head, freezing the embeddings and the
206
+ layers below (low-memory mode for GPUs under ~10 GB free). None trains everything."""
207
+ import torch
208
+ from laya.common import build_model, proper_reward
209
+ from safetensors.torch import load_file, save_file
210
+ from transformers import AutoTokenizer
211
+
212
+ records = read_jsonl(train) if isinstance(train, str) else list(train)
213
+ if not records:
214
+ raise JevError("sqljev: no training records")
215
+ device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
216
+ src = base_dir(base)
217
+ with open(os.path.join(src, "rl_agent_config.json")) as f:
218
+ cfg = json.load(f)
219
+ tok = AutoTokenizer.from_pretrained(os.path.join(src, "tokenizer"))
220
+ items, skipped = _items(records, tok, cfg)
221
+ if not items:
222
+ raise JevError("sqljev: none of the records could be encoded for this checkpoint")
223
+ rng = random.Random(seed)
224
+ rng.shuffle(items)
225
+ n_cal = min(400, len(items) // 10)
226
+ calib, items = items[:n_cal], items[n_cal:]
227
+ log("sqljev finetune: %d training items, %d held out for calibration%s, base %s, device %s"
228
+ % (len(items), len(calib), ", %d skipped" % skipped if skipped else "", base, device))
229
+
230
+ torch.manual_seed(seed)
231
+ model = build_model(cfg, encoder_dir=os.path.join(src, "encoder"))
232
+ model.load_state_dict(load_file(os.path.join(src, "model.safetensors")), strict=True)
233
+ if device.type == "cuda":
234
+ model.encoder.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False})
235
+ model.head_checkpointing = True
236
+ model.to(device).train()
237
+ if train_layers is not None:
238
+ n_layers = 1 + max(int(m.group(1)) for n, _ in model.named_parameters()
239
+ for m in [re.match(r"encoder\.layers\.(\d+)\.", n)] if m)
240
+ keep = n_layers - int(train_layers)
241
+ for n, p in model.named_parameters():
242
+ m = re.match(r"encoder\.layers\.(\d+)\.", n)
243
+ if n.startswith("encoder.embeddings.") or (m and int(m.group(1)) < keep):
244
+ p.requires_grad_(False)
245
+ log("sqljev finetune: training the top %d of %d encoder layers and the head" % (min(int(train_layers), n_layers), n_layers))
246
+
247
+ use_cuda = device.type == "cuda"
248
+ amp_dtype = torch.bfloat16 if use_cuda and torch.cuda.is_bf16_supported() else torch.float16
249
+ scaler = torch.amp.GradScaler("cuda", enabled=use_cuda and amp_dtype == torch.float16)
250
+ enc = [p for n, p in model.named_parameters() if n.startswith("encoder.") and p.requires_grad]
251
+ head = [p for n, p in model.named_parameters() if not n.startswith("encoder.") and p.requires_grad]
252
+ opt = torch.optim.AdamW([{"params": enc, "lr": lr_encoder}, {"params": head, "lr": lr_head}], weight_decay=0.01)
253
+ updates = max(1, -(-len(items) // (micro_batch * grad_accum)) * epochs)
254
+ sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=updates, eta_min=1e-6)
255
+ t0 = time.time()
256
+ for epoch in range(epochs):
257
+ rng.shuffle(items)
258
+ sigma = 0.4 + (0.1 - 0.4) * (epoch / max(1, epochs - 1))
259
+ total, steps = 0.0, 0
260
+ opt.zero_grad(set_to_none=True)
261
+ for b in range(0, len(items), micro_batch):
262
+ ids, att, mpos, mmask, target, qtype = (x.to(device) for x in _collate(items[b:b + micro_batch], tok.pad_token_id))
263
+ with torch.autocast(device.type, dtype=amp_dtype, enabled=use_cuda):
264
+ logits, act = model(ids, att, mpos, mmask, qtype)
265
+ logits = logits.float()
266
+ k = mmask.sum(-1, keepdim=True).float()
267
+ eps = torch.randn((group_size,) + logits.shape, device=device) * sigma * mmask
268
+ eps = (eps - eps.sum(-1, keepdim=True) / k) * mmask
269
+ z = logits.detach().unsqueeze(0) + eps
270
+ q = torch.softmax(z.masked_fill(~mmask, -1e4), -1)
271
+ with torch.no_grad():
272
+ r = proper_reward(q, target.unsqueeze(0), qtype, mmask, w_sph=0.75, w_rps=1.0)
273
+ adv = (r - r.mean(0, keepdim=True))
274
+ adv = adv / (adv.std() + 1e-6)
275
+ logp = -(((z - logits.unsqueeze(0)) ** 2) * mmask).sum(-1) / (2 * sigma ** 2)
276
+ loss_ce = -(target * torch.log_softmax(logits.masked_fill(~mmask, -1e4), -1)).sum(-1).mean()
277
+ loss = (-(adv * logp).mean() + loss_ce) / grad_accum + 0.0 * act.sum()
278
+ scaler.scale(loss).backward()
279
+ steps += 1
280
+ total += loss.item() * grad_accum
281
+ if steps % grad_accum == 0 or b + micro_batch >= len(items):
282
+ scaler.unscale_(opt)
283
+ torch.nn.utils.clip_grad_norm_(enc + head, 1.0)
284
+ scaler.step(opt)
285
+ scaler.update()
286
+ sched.step()
287
+ opt.zero_grad(set_to_none=True)
288
+ if steps % 100 == 0:
289
+ log(" epoch %d/%d step %d/%d loss %.4f %.0fs"
290
+ % (epoch + 1, epochs, steps, -(-len(items) // micro_batch), total / steps, time.time() - t0))
291
+ log(" epoch %d/%d done, mean loss %.4f, %.0fs" % (epoch + 1, epochs, total / max(1, steps), time.time() - t0))
292
+
293
+ # Calibrate one temperature per question type on items the run never trained on.
294
+ model.eval()
295
+ preds = []
296
+ with torch.no_grad():
297
+ for b in range(0, len(calib), 16):
298
+ chunk = calib[b:b + 16]
299
+ ids, att, mpos, mmask, target, qtype = (x.to(device) for x in _collate(chunk, tok.pad_token_id))
300
+ with torch.autocast(device.type, dtype=amp_dtype, enabled=use_cuda):
301
+ lg, _ = model(ids, att, mpos, mmask, qtype)
302
+ lg = lg.float().cpu().numpy()
303
+ preds += [(it["qtype"], lg[i, :len(it["markers"])].tolist(), it["target"]) for i, it in enumerate(chunk)]
304
+ temps = list(cfg.get("temperature") or [1.2, 1.2, 1.2])
305
+ for qt in range(3):
306
+ sel = [(z, t) for q, z, t in preds if q == qt]
307
+ if len(sel) >= 10:
308
+ temps[qt] = _fit_temperature(sel)
309
+
310
+ os.makedirs(out_dir, exist_ok=True)
311
+ save_file({k: v.half().contiguous().cpu() for k, v in model.state_dict().items()},
312
+ os.path.join(out_dir, "model.safetensors"))
313
+ model.encoder.config.save_pretrained(os.path.join(out_dir, "encoder"))
314
+ tok.save_pretrained(os.path.join(out_dir, "tokenizer"))
315
+ cfg.update(fine_tuned=True, model_name="laya-sqljev", temperature=temps)
316
+ cfg.pop("temperature_by_options", None)
317
+ with open(os.path.join(out_dir, "rl_agent_config.json"), "w") as f:
318
+ json.dump(cfg, f, indent=2)
319
+ questions = sorted({q["instructions"] for r in records for q in r["questions"].values()})
320
+ with open(os.path.join(out_dir, "sqljev_finetune.json"), "w") as f:
321
+ json.dump({"base": base, "records": len(records), "items": len(items), "epochs": epochs,
322
+ "questions": questions, "seconds": round(time.time() - t0, 1)}, f, indent=2)
323
+ log("sqljev finetune: saved to %s (%.0fs)" % (out_dir, time.time() - t0))
324
+ return out_dir
325
+
326
+
327
+ # ---------------------------------------------------------------- publishing
328
+
329
+ def publish(out_dir, repo_id, token=None, private=True, metrics=None):
330
+ """Upload a fine-tuned checkpoint to the Hugging Face Hub with a model card. Every database then uses it
331
+ with SQLJEV_MODEL=<repo_id>."""
332
+ from huggingface_hub import HfApi
333
+ meta_path = os.path.join(out_dir, "sqljev_finetune.json")
334
+ meta = json.load(open(meta_path)) if os.path.exists(meta_path) else {}
335
+ if metrics:
336
+ meta["metrics"] = metrics
337
+ with open(meta_path, "w") as f:
338
+ json.dump(meta, f, indent=2)
339
+ card = """---
340
+ license: apache-2.0
341
+ base_model: %s
342
+ tags: [laya, sqljev, sql, decision-model]
343
+ ---
344
+
345
+ # %s
346
+
347
+ A [Laya](https://github.com/NandhaKishorM/laya) checkpoint fine-tuned with [sqljev](https://github.com/singhpratech/sqljev)
348
+ on labelled SQL rows.
349
+
350
+ Questions it was trained on:
351
+ %s
352
+
353
+ %s
354
+ Use it from any database sqljev supports:
355
+
356
+ ```bash
357
+ SQLJEV_MODEL=%s sqljev gateway --host 0.0.0.0
358
+ ```
359
+ """ % (meta.get("base", LAYA_BASE), repo_id.split("/")[-1],
360
+ "\n".join("- %s" % q for q in meta.get("questions", [])),
361
+ ("Held-out accuracy: %s\n" % json.dumps(metrics)) if metrics else "", repo_id)
362
+ with open(os.path.join(out_dir, "README.md"), "w") as f:
363
+ f.write(card)
364
+ api = HfApi(token=token)
365
+ api.create_repo(repo_id, private=private, exist_ok=True)
366
+ api.upload_folder(repo_id=repo_id, folder_path=out_dir, commit_message="sqljev finetune")
367
+ return "https://huggingface.co/" + repo_id
sqljev/gateway.py ADDED
@@ -0,0 +1,215 @@
1
+ """sqljev gateway: one HTTP service that answers the batch protocols databases use to call out.
2
+
3
+ sqljev gateway --port 8765 # Laya in-process (backend 'local') by default
4
+
5
+ Routes (all POST, JSON):
6
+
7
+ /v1/eval sqljev native: {"question", "kind", "options", "rows": [...]} -> {"answers": [...]}
8
+ /sqlserver same body; called by jev.judge through sp_invoke_external_rest_endpoint
9
+ /bigquery BigQuery remote functions: {"calls": [[args]], "userDefinedContext": {"fn": ...}}
10
+ -> {"replies": [...]}
11
+ /snowflake/<fn> Snowflake external / service functions: {"data": [[rownum, args...]]}
12
+ -> {"data": [[rownum, result]]}
13
+ /redshift Redshift Lambda UDF payload over HTTP (the same handler as sqljev.aws_lambda)
14
+ /v1/systemone TypeSafe Jev's wire protocol, for Jev clients such as pg-jev: batch requests
15
+ {"state": {"condition", "rows": [...]}, "questions": {"r0": ...}} are answered by the
16
+ engine, so `SET jev.api_url = 'http://gateway:8765/v1/systemone'` runs pg-jev on Laya
17
+ GET /health, GET /stats
18
+
19
+ Each database already sends rows in batches (BigQuery up to max_batching_rows, Snowflake/Redshift in their
20
+ own batch sizes), so every request becomes one engine.call(): rows grouped by question, de-duplicated,
21
+ looked up in the cache and the misses judged in shared forward passes.
22
+
23
+ Auth: set SQLJEV_GATEWAY_TOKEN and send it as "Authorization: Bearer <token>" or "X-Jev-Token: <token>".
24
+ BigQuery cannot send headers: run the gateway on Cloud Run with IAM auth instead.
25
+ """
26
+ import hmac
27
+ import json
28
+ import os
29
+ import re
30
+ import sys
31
+ import traceback
32
+ from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
33
+
34
+ from .core import FUNCTIONS, JevError, Jev, __version__, function_name
35
+
36
+ MAX_BODY_BYTES = 64 * 1024 * 1024
37
+
38
+
39
+ def handle_eval(engine, body):
40
+ kind = body.get("kind") or "noul"
41
+ rows = body.get("rows")
42
+ if not isinstance(rows, list) or not body.get("question"):
43
+ raise JevError("sqljev: body needs 'question' and a 'rows' array")
44
+ return {"answers": engine.evaluate(rows, body["question"], kind, body.get("options"))}
45
+
46
+
47
+ def handle_bigquery(engine, body):
48
+ ctx = body.get("userDefinedContext") or {}
49
+ fn = ctx.get("fn") or ctx.get("function")
50
+ if not fn:
51
+ raise JevError("sqljev: set user_defined_context = [('fn', 'jev_prob')] on the remote function")
52
+ return {"replies": _json_safe(engine.call(fn, body.get("calls") or []))}
53
+
54
+
55
+ def handle_snowflake(engine, body, fn):
56
+ data = body.get("data") or []
57
+ results = engine.call(fn, [r[1:] for r in data]) # VARIANT takes jev_eval's object as is
58
+ return {"data": [[r[0], v] for r, v in zip(data, results)]}
59
+
60
+
61
+ def handle_redshift(engine, event):
62
+ """Redshift Lambda UDF protocol. Errors are reported in-band, as Redshift expects."""
63
+ try:
64
+ fn = event.get("external_function") or ""
65
+ args = event.get("arguments") or []
66
+ results = engine.call(fn, args)
67
+ return {"success": True, "num_records": len(results), "results": _json_safe(results)}
68
+ except Exception as e: # noqa: BLE001 -- Redshift shows error_msg to the user
69
+ return {"success": False, "error_msg": str(e)}
70
+
71
+
72
+ _ROW_REF = re.compile(r"`rows\[(\d+)\]`")
73
+ _JEV_PREFIX = {"score": "Rate the record `rows[%d]`: ", "choice": "For the record `rows[%d]`: "}
74
+
75
+
76
+ def handle_systemone(engine, body):
77
+ """Answer a Jev batch request (one shared state with `rows`, one question per row, the format pg-jev and
78
+ sqljev's own jev backend send) with the engine: questions are grouped back into (question, rows) pairs."""
79
+ state, questions = body.get("state"), body.get("questions")
80
+ if not isinstance(state, dict) or not isinstance(state.get("rows"), list) or not isinstance(questions, dict):
81
+ raise JevError("sqljev: /v1/systemone accepts batch requests: state {'rows': [...]} with one question per "
82
+ "row that names `rows[i]` (as pg-jev sends); use /v1/eval for everything else")
83
+ rows, groups = state["rows"], {}
84
+ for qid, q in questions.items():
85
+ m = _ROW_REF.search(q.get("instructions") or "")
86
+ if not m or int(m.group(1)) >= len(rows):
87
+ raise JevError("sqljev: question %r does not name a row as `rows[i]`" % qid)
88
+ i, kind = int(m.group(1)), q.get("type")
89
+ if kind == "noul":
90
+ query, opts = state.get("condition") or "", None
91
+ elif kind in _JEV_PREFIX:
92
+ prefix = _JEV_PREFIX[kind] % i
93
+ instr = q["instructions"]
94
+ query = instr[len(prefix):] if instr.startswith(prefix) else instr
95
+ crit = q.get("criteria")
96
+ opts = list(crit.keys()) if isinstance(crit, dict) else list(crit or [])
97
+ else:
98
+ raise JevError("sqljev: unknown question type %r" % kind)
99
+ groups.setdefault((kind, query, tuple(opts or ())), []).append((qid, i))
100
+ answers = {}
101
+ for (kind, query, opts), items in groups.items():
102
+ got = engine.evaluate([rows[i] for _, i in items], query, kind, list(opts) or None)
103
+ answers.update({qid: a for (qid, _), a in zip(items, got)})
104
+ return {"model": "sqljev-" + engine.cfg["backend"], "answers": answers,
105
+ "usage": {"input_tokens": 0, "output_tokens": 0}}
106
+
107
+
108
+ def _json_safe(values):
109
+ # jev_eval answers go back to warehouses as JSON text columns (BigQuery JSON/STRING, Snowflake VARIANT
110
+ # accepts objects directly but a string is safe everywhere).
111
+ return [json.dumps(v) if isinstance(v, dict) else v for v in values]
112
+
113
+
114
+ class Handler(BaseHTTPRequestHandler):
115
+ protocol_version = "HTTP/1.1"
116
+ engine = None
117
+ token = None
118
+ server_version = "sqljev/" + __version__
119
+
120
+ def log_message(self, fmt, *args):
121
+ if os.environ.get("SQLJEV_GATEWAY_LOG"):
122
+ sys.stderr.write("sqljev gateway: " + fmt % args + "\n")
123
+
124
+ def _send(self, code, obj):
125
+ data = json.dumps(obj).encode()
126
+ self.send_response(code)
127
+ self.send_header("Content-Type", "application/json")
128
+ self.send_header("Content-Length", str(len(data)))
129
+ self.end_headers()
130
+ self.wfile.write(data)
131
+
132
+ def _authorized(self):
133
+ if not self.token:
134
+ return True
135
+ got = self.headers.get("X-Jev-Token") or ""
136
+ auth = self.headers.get("Authorization") or ""
137
+ if auth.startswith("Bearer "):
138
+ got = got or auth[7:]
139
+ return hmac.compare_digest(got.encode(), self.token.encode())
140
+
141
+ def do_GET(self):
142
+ if self.path == "/health":
143
+ return self._send(200, {"status": "ok", "version": __version__, "backend": self.engine.cfg["backend"],
144
+ "functions": sorted(FUNCTIONS)})
145
+ if self.path == "/stats":
146
+ if not self._authorized():
147
+ return self._send(401, {"error": "unauthorized"})
148
+ return self._send(200, self.engine.stats())
149
+ self._send(404, {"error": "not found"})
150
+
151
+ def do_POST(self):
152
+ n = int(self.headers.get("Content-Length") or 0)
153
+ if n > MAX_BODY_BYTES:
154
+ return self._send(413, {"error": "body too large"})
155
+ raw = self.rfile.read(n)
156
+ if not self._authorized():
157
+ return self._send(401, {"error": "unauthorized"})
158
+ path = self.path.split("?", 1)[0].rstrip("/")
159
+ try:
160
+ body = json.loads(raw or b"{}")
161
+ if path in ("/v1/eval", "/sqlserver"):
162
+ return self._send(200, handle_eval(self.engine, body))
163
+ if path == "/bigquery":
164
+ return self._send(200, handle_bigquery(self.engine, body))
165
+ if path.startswith("/snowflake/"):
166
+ return self._send(200, handle_snowflake(self.engine, body, function_name(path.rsplit("/", 1)[1])))
167
+ if path == "/redshift":
168
+ return self._send(200, handle_redshift(self.engine, body))
169
+ if path == "/v1/systemone":
170
+ return self._send(200, handle_systemone(self.engine, body))
171
+ return self._send(404, {"error": "not found: " + path})
172
+ except (JevError, ValueError, KeyError, TypeError, IndexError) as e:
173
+ # 400 is not retried by BigQuery/Snowflake: validation errors surface as query errors.
174
+ return self._send(400, {"error": str(e), "errorMessage": str(e)})
175
+ except Exception as e: # noqa: BLE001
176
+ traceback.print_exc()
177
+ return self._send(500, {"error": "internal error: %s" % type(e).__name__,
178
+ "errorMessage": "sqljev gateway internal error"})
179
+
180
+
181
+ def make_server(host="127.0.0.1", port=8765, engine=None, token=None):
182
+ handler = type("SqlJevHandler", (Handler,), {
183
+ "engine": engine or Jev(),
184
+ "token": token if token is not None else os.environ.get("SQLJEV_GATEWAY_TOKEN"),
185
+ })
186
+ ThreadingHTTPServer.daemon_threads = True
187
+ return ThreadingHTTPServer((host, port), handler)
188
+
189
+
190
+ def serve(host="127.0.0.1", port=8765, engine=None, token=None, certfile=None, keyfile=None):
191
+ """certfile/keyfile serve HTTPS directly (sp_invoke_external_rest_endpoint only calls https:// URLs);
192
+ behind a TLS-terminating proxy or Cloud Run, leave them unset."""
193
+ srv = make_server(host, port, engine, token)
194
+ scheme = "http"
195
+ if certfile:
196
+ import ssl
197
+ ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
198
+ ctx.load_cert_chain(certfile, keyfile)
199
+ accept = srv.get_request
200
+
201
+ def get_request():
202
+ # Handshake lazily, in the worker thread: a slow or failing client never blocks accept().
203
+ sock, addr = accept()
204
+ return ctx.wrap_socket(sock, server_side=True, do_handshake_on_connect=False), addr
205
+ srv.get_request = get_request
206
+ scheme = "https"
207
+ eng = srv.RequestHandlerClass.engine
208
+ print("sqljev gateway %s on %s://%s:%d (backend %s)" % (__version__, scheme, host, port, eng.cfg["backend"]),
209
+ file=sys.stderr)
210
+ try:
211
+ srv.serve_forever()
212
+ except KeyboardInterrupt:
213
+ pass
214
+ finally:
215
+ srv.server_close()
sqljev/spark.py ADDED
@@ -0,0 +1,64 @@
1
+ """Spark / Databricks: the jev functions as pandas (Arrow) UDFs usable from SQL.
2
+
3
+ import sqljev.spark
4
+ sqljev.spark.register(spark) # Laya in-process on each executor
5
+ # or: sqljev.spark.register(spark, backend="gateway", api_url="https://.../v1/eval",
6
+ # api_key=dbutils.secrets.get("sqljev", "token"))
7
+
8
+ SELECT * FROM tickets WHERE jev(to_json(struct(*)), 'the customer is angry')
9
+
10
+ Spark passes Arrow batches (spark.sql.execution.arrow.maxRecordsPerBatch, 10,000 rows by default) to each
11
+ UDF; each batch becomes one engine.call(). One engine per Python worker keeps the model loaded and the
12
+ answer cache warm across batches. On a GPU cluster set device="cuda".
13
+ """
14
+ _ENGINES = {}
15
+
16
+ _RETURNS = {"jev": "boolean", "jev_prob": "double", "jev_score": "double", "jev_score_norm": "double",
17
+ "jev_choice": "string", "jev_confidence": "double", "jev_eval": "string"}
18
+ _ARITY = {"jev": 2, "jev_prob": 2, "jev_score": 3, "jev_score_norm": 3, "jev_choice": 3,
19
+ "jev_confidence": 4, "jev_eval": 4}
20
+
21
+
22
+ def _engine(settings):
23
+ from .core import Jev
24
+ key = tuple(sorted(settings.items()))
25
+ eng = _ENGINES.get(key)
26
+ if eng is None:
27
+ eng = _ENGINES[key] = Jev(**settings)
28
+ return eng
29
+
30
+
31
+ def _make(name, settings):
32
+ import json
33
+ import pandas as pd
34
+ from pyspark.sql.functions import pandas_udf
35
+
36
+ n = _ARITY[name]
37
+
38
+ def run(*cols):
39
+ calls = list(zip(*(c.tolist() for c in cols)))
40
+ out = _engine(settings).call(name, calls)
41
+ if name == "jev_eval":
42
+ out = [None if v is None else json.dumps(v) for v in out]
43
+ return pd.Series(out, dtype="object" if name in ("jev_choice", "jev_eval", "jev") else "float64")
44
+
45
+ # pandas_udf needs a fixed-arity signature with type hints.
46
+ if n == 2:
47
+ def f(a: pd.Series, b: pd.Series) -> pd.Series:
48
+ return run(a, b)
49
+ elif n == 3:
50
+ def f(a: pd.Series, b: pd.Series, c: pd.Series) -> pd.Series:
51
+ return run(a, b, c)
52
+ else:
53
+ def f(a: pd.Series, b: pd.Series, c: pd.Series, d: pd.Series) -> pd.Series:
54
+ return run(a, b, c, d)
55
+ return pandas_udf(f, _RETURNS[name])
56
+
57
+
58
+ def register(spark, prefix="", **settings):
59
+ """Register jev, jev_prob, jev_score, jev_score_norm, jev_choice, jev_confidence, jev_eval for SQL.
60
+ Settings (backend, api_url, api_key, model, device, ...) are shipped to the executors."""
61
+ fns = {}
62
+ for name in _RETURNS:
63
+ fns[name] = spark.udf.register(prefix + name, _make(name, dict(settings)))
64
+ return fns