128 lines
5.4 KiB
Python
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)
|