Setting the file. One moment. Snapshot SQLite State · Eval Engineering · langchain-ai/langchain-skills · Skills DocsScript
scripts/snapshot_sqlite_state.py
Python·156 lines·5 KB
from
typing
import
Any, Optional
14
15
16class SnapshotError(ValueError):
17 """Raised when snapshot input is invalid."""
18
19
20DENIED_ACTIONS = {
21 sqlite3.SQLITE_ATTACH,
22 sqlite3.SQLITE_ALTER_TABLE,
23 sqlite3.SQLITE_CREATE_INDEX,
24 sqlite3.SQLITE_CREATE_TABLE,
25 sqlite3.SQLITE_CREATE_TEMP_INDEX,
26 sqlite3.SQLITE_CREATE_TEMP_TABLE,
27 sqlite3.SQLITE_CREATE_TEMP_TRIGGER,
28 sqlite3.SQLITE_CREATE_TEMP_VIEW,
29 sqlite3.SQLITE_CREATE_TRIGGER,
30 sqlite3.SQLITE_CREATE_VIEW,
31 sqlite3.SQLITE_DELETE,
32 sqlite3.SQLITE_DETACH,
33 sqlite3.SQLITE_DROP_INDEX,
34 sqlite3.SQLITE_DROP_TABLE,
35 sqlite3.SQLITE_DROP_TEMP_INDEX,
36 sqlite3.SQLITE_DROP_TEMP_TABLE,
37 sqlite3.SQLITE_DROP_TEMP_TRIGGER,
38 sqlite3.SQLITE_DROP_TEMP_VIEW,
39 sqlite3.SQLITE_DROP_TRIGGER,
40 sqlite3.SQLITE_DROP_VIEW,
41 sqlite3.SQLITE_INSERT,
42 sqlite3.SQLITE_PRAGMA,
43 sqlite3.SQLITE_REINDEX,
44 sqlite3.SQLITE_TRANSACTION,
45 sqlite3.SQLITE_UPDATE,
46}
47
48
49def load_config(path: Path) -> list[dict[str, Any]]:
50 try:
51 value = json.loads(path.read_text(encoding="utf-8"))
52 except (OSError, json.JSONDecodeError) as error:
53 raise SnapshotError(f"cannot read {path}: {error}") from error
54 queries = value.get("queries") if isinstance(value, dict) else None
55 if not isinstance(queries, list) or not queries:
56 raise SnapshotError("config must contain a non-empty queries list")
57
58 names: set[str] = set()
59 for query in queries:
60 if not isinstance(query, dict):
61 raise SnapshotError("each query must be an object")
62 name = query.get("name")
63 sql = query.get("sql")
64 params = query.get("params", [])
65 preserve_order = query.get("preserve_order", False)
66 if not isinstance(name, str) or not name or name in names:
67 raise SnapshotError("query names must be unique non-empty strings")
68 if not isinstance(sql, str) or not sql.lstrip().lower().startswith(("select ", "with ")):
69 raise SnapshotError(f"query {name} must be a SELECT or WITH statement")
70 if ";" in sql.rstrip().rstrip(";"):
71 raise SnapshotError(f"query {name} must contain one statement")
72 if not isinstance(params, list):
73 raise SnapshotError(f"query {name} params must be a list")
74 if type(preserve_order) is not bool:
75 raise SnapshotError(f"query {name} preserve_order must be boolean")
76 names.add(name)
77 return queries
78
79
80def authorizer(
81 action: int,
82 _arg1: Optional[str],
83 _arg2: Optional[str],
84 _db: Optional[str],
85 _source: Optional[str],
86) -> int:
87 return sqlite3.SQLITE_DENY if action in DENIED_ACTIONS else sqlite3.SQLITE_OK
88
89
90def json_value(value: Any) -> Any:
91 if isinstance(value, bytes):
92 return {"$type": "bytes", "base64": base64.b64encode(value).decode("ascii")}
93 if isinstance(value, float) and not math.isfinite(value):
94 return {"$type": "float", "value": repr(value)}
95 return value
96
97
98def canonical_rows(cursor: sqlite3.Cursor, preserve_order: bool = False) -> list[dict[str, Any]]:
99 columns = [item[0] for item in cursor.description or []]
100 if len(columns) != len(set(columns)):
101 raise SnapshotError("query result contains duplicate column names")
102 rows = [
103 {column: json_value(value) for column, value in zip(columns, row)}
104 for row in cursor.fetchall()
105 ]
106 if preserve_order:
107 return rows
108 return sorted(rows, key=lambda row: json.dumps(row, sort_keys=True, separators=(",", ":")))
109
110
111def capture(database: Path, config: Path) -> dict[str, Any]:
112 if not database.is_file():
113 raise SnapshotError(f"database not found: {database}")
114 queries = load_config(config)
115 uri = f"{database.resolve().as_uri()}?mode=ro"
116 connection = sqlite3.connect(uri, uri=True)
117 connection.set_authorizer(authorizer)
118 try:
119 results: dict[str, Any] = {}
120 for query in queries:
121 try:
122 cursor = connection.execute(query["sql"], query.get("params", []))
123 except sqlite3.Error as error:
124 raise SnapshotError(f"query {query['name']} failed: {error}") from error
125 results[query["name"]] = canonical_rows(
126 cursor, query.get("preserve_order", False)
127 )
128 finally:
129 connection.close()
130
131 return {"results": results}
132
133
134def main() -> int:
135 parser = argparse.ArgumentParser(description=__doc__)
136 parser.add_argument("--db", required=True, type=Path)
137 parser.add_argument("--config", required=True, type=Path)
138 parser.add_argument("--output", required=True, type=Path)
139 args = parser.parse_args()
140
141 try:
142 snapshot = capture(args.db, args.config)
143 args.output.parent.mkdir(parents=True, exist_ok=True)
144 args.output.write_text(
145 json.dumps(snapshot, indent=2, sort_keys=True) + "\n",
146 encoding="utf-8",
147 )
148 except (OSError, SnapshotError, sqlite3.Error) as error:
149 print(f"ERROR: {error}", file=sys.stderr)
150 return 2
151 print(args.output)
152 return 0
153
154
155if __name__ == "__main__":
156 raise SystemExit(main())