-
1
-
2
-
3
-
4
-
5
-
6
-
7
-
8
-
9
-
10
-
11
-
12
-
13
-
14
-
15
-
16
-
17
-
18
-
19
-
20
-
21
-
22
-
23
-
24
-
25
-
26
-
27
-
28
-
29
-
30
-
31
-
32
-
33
-
34
-
35
-
36
-
37
-
38
-
39
-
40
-
41
-
42
-
43
-
44
-
45
-
46
-
47
-
48
-
49
-
50
-
51
-
52
-
53
-
54
-
55
-
56
-
57
-
58
-
59
-
60
-
61
-
62
-
63
-
64
-
65
-
66
-
67
-
68
-
69
-
70
-
71
-
72
-
73
-
74
-
75
-
76
-
77
-
78
-
79
-
80
-
81
-
82
-
83
-
84
-
85
-
86
-
87
-
88
-
89
-
90
-
91
-
92
-
93
-
94
-
95
-
96
-
97
-
98
-
99
-
100
-
101
-
102
-
103
-
104
-
105
-
106
-
107
-
108
-
109
-
110
-
111
-
112
-
113
-
114
-
115
-
116
-
117
-
118
-
119
-
120
-
121
-
122
-
123
-
124
-
125
-
126
-
127
-
128
-
129
-
130
-
131
-
132
-
133
-
134
-
135
-
136
-
137
-
138
-
139
-
140
-
141
-
142
-
143
-
144
-
145
-
146
-
147
-
148
-
149
-
150
-
151
-
152
-
153
-
154
-
155
-
156
-
157
-
158
-
159
-
160
-
161
-
162
-
163
-
164
-
165
-
166
-
167
-
168
-
169
-
170
-
171
-
172
-
173
-
174
-
175
-
176
import mimetypes
import os
import pathlib
import sqlite3
import sys
import threading
import numpy as np
import PIL
import sqlite_vec
from watchdog.events import FileOpenedEvent, FileSystemEventHandler
from watchdog.observers import Observer
import model
from unixsocket import UnixStreamXMLRPCServer
DIM = 768
print("Connecting to DB")
con = sqlite3.connect("search.db")
con.execute("PRAGMA journal_mode=WAL")
con.enable_load_extension(True)
sqlite_vec.load(con)
con.enable_load_extension(False)
cur = con.cursor()
cur.execute(
f"CREATE VIRTUAL TABLE IF NOT EXISTS idx USING vec0(ino INTEGER PRIMARY KEY, emb float[{DIM}], time INTEGER, parent INTEGER, path TEXT)",
)
con.commit()
lock = threading.Lock()
def get_parent(path):
# Get inode of parent
if path in watchdirs:
return 0
parent = pathlib.Path(path).parent
return os.stat(parent).st_ino
class EventHandler(FileSystemEventHandler):
def dispatch(self, event):
if not isinstance(event, FileOpenedEvent):
with lock:
# print(event)
super().dispatch(event)
def on_created(self, event):
index(event.src_path, get_parent(event.src_path))
def on_modified(self, event):
if not event.is_directory:
self.on_created(event)
def on_deleted(self, event):
res = cur.execute("SELECT ino FROM idx WHERE path = ?", (event.src_path,))
unindex(res.fetchone()[0])
def on_moved(self, event):
# inode doesn't change after move
print("Moving", event.src_path, event.dest_path)
s = os.stat(event.dest_path)
cur.execute(
"UPDATE idx SET time = ?, parent = ?, path = ? WHERE ino = ?",
(s.st_mtime_ns, get_parent(event.dest_path), event.dest_path, s.st_ino),
)
cur.execute(
"UPDATE idx SET path = replace(path, ?, ?)",
(event.src_path, event.dest_path),
)
con.commit()
def index(path, parent):
if not os.path.exists(path) or os.path.basename(path).startswith("."):
# Skip nonexistent or hidden files
return
print("Indexing", path, parent)
s = os.stat(path)
res = cur.execute("SELECT time, parent, path FROM idx WHERE ino = ?", (s.st_ino,))
db_vals = res.fetchone()
if os.path.isfile(path):
if db_vals is None or db_vals[0] != s.st_mtime_ns or db_vals[1] != parent:
# Not in DB or modified
# Probably faster to query DB first instead of guessing mimetype first
mtype = mimetypes.guess_type(path)[0]
if mtype is None or not mtype.startswith("image"):
# Only support image embeddings for now
return
try:
emb = model.embed_image(path)
except PIL.UnidentifiedImageError:
print("Couldn't index", path)
return
# sqlite-vec doesn't support INSERT OR REPLACE and UPSERT
# https://github.com/asg017/sqlite-vec/issues/127
if db_vals is not None:
cur.execute("DELETE FROM idx WHERE ino = ?", (s.st_ino,))
cur.execute(
"INSERT INTO idx VALUES (?, ?, ?, ?, ?)",
(s.st_ino, emb, s.st_mtime_ns, parent, path),
)
con.commit()
elif db_vals[2] != path:
# Moved
cur.execute("UPDATE idx SET path = ? WHERE ino = ?", (path, s.st_ino))
con.commit()
elif os.path.isdir(path):
if db_vals != (s.st_mtime_ns, parent, path):
if db_vals is not None:
cur.execute("DELETE FROM idx WHERE ino = ?", (s.st_ino,))
emb = np.ones((DIM,), dtype=np.float32)
cur.execute(
"INSERT INTO idx VALUES (?, ?, ?, ?, ?)",
(s.st_ino, emb, s.st_mtime_ns, parent, path),
)
con.commit()
for child in os.listdir(path):
index(os.path.join(path, child), s.st_ino)
def unindex(ino):
print("Unindexing", ino)
res = cur.execute("SELECT id FROM idx WHERE parent = ?", (ino,))
for db_child in res.fetchall():
unindex(db_child[0])
cur.execute("DELETE FROM idx WHERE ino = ?", (ino,))
con.commit()
def search(text, limit):
print("Search", text, limit)
if os.path.exists(text):
emb = model.embed_image(text)
else:
emb = model.embed_text(text)
res = cur.execute(
"SELECT path FROM idx WHERE emb MATCH ? AND k = ? ORDER BY distance",
(emb, limit),
)
return [i[0] for i in res.fetchall()]
print("Indexing files")
watchdirs = set(map(os.path.abspath, sys.argv[1:]))
observer = Observer()
observer.start()
event_handler = EventHandler()
for wdir in watchdirs:
observer.schedule(event_handler, wdir, recursive=True)
with lock:
for wdir in watchdirs:
# Pretend 0 is parent of watchdirs
index(wdir, 0)
# Remove stale entries
res = cur.execute("SELECT ino, path FROM idx")
for ino, path in res.fetchall():
if not os.path.exists(path):
cur.execute("DELETE FROM idx WHERE ino = ?", (ino,))
con.commit()
print("Starting RPC server")
sockpath = os.path.join(os.environ["XDG_RUNTIME_DIR"], "search.sock")
pathlib.Path(sockpath).unlink(missing_ok=True)
server = UnixStreamXMLRPCServer(sockpath)
server.register_function(search)
server.serve_forever()