Files

128 lines
5.4 KiB
Python

import json
import sqlite3
import uuid
from contextlib import contextmanager
from datetime import datetime, timezone
def now():
return datetime.now(timezone.utc).isoformat()
class Conflict(Exception):
pass
class Store:
def __init__(self, path):
self.path = str(path)
with self.db() as db:
db.executescript('''
PRAGMA journal_mode=WAL;
CREATE TABLE IF NOT EXISTS prompts (
id TEXT PRIMARY KEY, data TEXT NOT NULL, version INTEGER NOT NULL,
created TEXT NOT NULL, updated TEXT NOT NULL, deleted INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE IF NOT EXISTS versions (
prompt_id TEXT NOT NULL REFERENCES prompts(id), version INTEGER NOT NULL,
data TEXT NOT NULL, note TEXT NOT NULL, created TEXT NOT NULL,
PRIMARY KEY(prompt_id, version)
);
CREATE TABLE IF NOT EXISTS vectors (
prompt_id TEXT PRIMARY KEY REFERENCES prompts(id), version INTEGER NOT NULL,
fingerprint TEXT NOT NULL, vector TEXT NOT NULL
);
''')
@contextmanager
def db(self):
db = sqlite3.connect(self.path, timeout=10)
db.row_factory = sqlite3.Row
db.execute('PRAGMA foreign_keys=ON')
try:
with db:
yield db
finally:
db.close()
@staticmethod
def decode(row):
return dict(json.loads(row['data']), id=row['id'], version=row['version'],
created=row['created'], updated=row['updated'], deleted=bool(row['deleted']))
def list(self, deleted=False):
with self.db() as db:
return [self.decode(r) for r in db.execute(
'SELECT * FROM prompts WHERE deleted=? ORDER BY updated DESC', (int(deleted),))]
def get(self, ident):
with self.db() as db:
row = db.execute('SELECT * FROM prompts WHERE id=?', (ident,)).fetchone()
return self.decode(row) if row else None
def save(self, data, ident=None, expected=None, note='Gespeichert'):
stamp = now()
with self.db() as db:
db.execute('BEGIN IMMEDIATE')
if ident:
row = db.execute('SELECT * FROM prompts WHERE id=?', (ident,)).fetchone()
if not row:
raise KeyError(ident)
if row['version'] != expected:
raise Conflict('Der Prompt wurde inzwischen geändert. Bitte neu öffnen.')
version = expected + 1
db.execute('UPDATE prompts SET data=?,version=?,updated=? WHERE id=?',
(json.dumps(data), version, stamp, ident))
db.execute('DELETE FROM vectors WHERE prompt_id=?', (ident,))
else:
ident, version = str(uuid.uuid4()), 1
db.execute('INSERT INTO prompts VALUES (?,?,?,?,?,0)',
(ident, json.dumps(data), version, stamp, stamp))
db.execute('INSERT INTO versions VALUES (?,?,?,?,?)',
(ident, version, json.dumps(data), note, stamp))
return self.get(ident)
def trash(self, ident, deleted):
with self.db() as db:
result = db.execute('UPDATE prompts SET deleted=?,updated=? WHERE id=?',
(int(deleted), now(), ident))
if not result.rowcount:
raise KeyError(ident)
def history(self, ident):
with self.db() as db:
return [dict(version=r['version'], data=json.loads(r['data']), note=r['note'],
created=r['created']) for r in db.execute(
'SELECT * FROM versions WHERE prompt_id=? ORDER BY version DESC', (ident,))]
def put_vector(self, ident, version, fingerprint, vector):
with self.db() as db:
db.execute('''INSERT OR REPLACE INTO vectors
SELECT id,version,?,? FROM prompts WHERE id=? AND version=? AND deleted=0''',
(fingerprint, json.dumps(vector), ident, version))
def vectors(self, fingerprint):
with self.db() as db:
return {r['prompt_id']: json.loads(r['vector']) for r in db.execute('''
SELECT v.* FROM vectors v JOIN prompts p ON p.id=v.prompt_id
WHERE v.fingerprint=? AND v.version=p.version AND p.deleted=0''', (fingerprint,))}
def export(self):
return {'format': 'casadeprompt', 'schema': 1, 'exported': now(),
'prompts': [dict(p, versions=self.history(p['id']))
for p in self.list() + self.list(True)]}
def import_items(self, items):
# Caller validates the complete archive before this transaction begins.
with self.db() as db:
for item in items:
ident, stamp = str(uuid.uuid4()), now()
versions = item['versions']
db.execute('INSERT INTO prompts VALUES (?,?,?,?,?,?)',
(ident, json.dumps(item['data']), len(versions), stamp, stamp,
int(item.get('deleted', False))))
for i, entry in enumerate(versions, 1):
db.execute('INSERT INTO versions VALUES (?,?,?,?,?)',
(ident, i, json.dumps(entry['data']), entry['note'], entry['created']))
return len(items)