forked from BasedHardware/omi
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathspeaker_sample_migration.py
More file actions
299 lines (241 loc) · 11.3 KB
/
Copy pathspeaker_sample_migration.py
File metadata and controls
299 lines (241 loc) · 11.3 KB
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
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
"""
Speaker sample migration utility.
Provides functions for migrating speaker samples across versions:
- v1 → v2: Add transcripts to samples
- v2 → v3: Regenerate embeddings using /v2/embedding API
- v1 → v3: Full migration (transcripts + new embeddings)
Uses in-process locking to prevent concurrent migrations.
"""
from typing import Any, Dict, List
import asyncio
from google.cloud.exceptions import NotFound
from database import users as users_db
from utils.speaker_sample import (
delete_sample_from_storage,
download_sample_audio,
verify_and_transcribe_sample,
)
from utils.executors import db_executor, storage_executor, sync_executor, run_blocking
from utils.stt.speaker_embedding import extract_embedding_from_bytes
import logging
logger = logging.getLogger(__name__)
# In-process locks to prevent concurrent migration for same person
_migration_locks: dict[tuple[str, str], asyncio.Lock] = {}
_locks_lock = asyncio.Lock()
async def _get_migration_lock(uid: str, person_id: str) -> asyncio.Lock:
"""Get or create a lock for the given uid/person_id pair."""
key = (uid, person_id)
async with _locks_lock:
if key not in _migration_locks:
_migration_locks[key] = asyncio.Lock()
return _migration_locks[key]
async def migrate_person_samples_v1_to_v2(uid: str, person: Dict[str, Any]) -> Dict[str, Any]:
"""
Migrate person's speech samples from v1 to v2.
v1: Only speech_samples (paths), no transcripts
v2: speech_samples + speech_sample_transcripts (parallel arrays)
Samples that fail quality checks are DROPPED along with speaker_embedding.
Uses in-process lock to prevent concurrent migration for same person.
Args:
uid: User ID
person: Person dict with 'id', 'speech_samples', 'speech_samples_version', etc.
Returns:
Updated person dict with migrated fields
"""
version = person.get('speech_samples_version', 1)
if version >= 2:
return person
person_id = person['id']
lock = await _get_migration_lock(uid, person_id)
async with lock:
# Re-check version inside lock (another call may have migrated)
fresh_person = await run_blocking(db_executor, users_db.get_person, uid, person_id)
if fresh_person and fresh_person.get('speech_samples_version', 1) >= 2:
return fresh_person
samples = person.get('speech_samples', [])
if not samples:
if person.get('speaker_embedding'):
await run_blocking(db_executor, users_db.clear_person_speaker_embedding, uid, person_id)
logger.info(f"v1→v2 migration: cleared stale embedding for person with no samples {uid} {person_id}")
await run_blocking(db_executor, users_db.update_person_speech_samples_version, uid, person_id, 2)
person['speech_samples_version'] = 2
person['speech_sample_transcripts'] = []
person['speaker_embedding'] = None
return person
valid_samples: List[Any] = []
valid_transcripts: List[Any] = []
samples_to_delete: List[Any] = []
has_transient_failures = False
for sample_path in samples:
try:
audio_bytes = await run_blocking(storage_executor, download_sample_audio, sample_path)
except NotFound:
logger.warning(f"Sample not found in storage, skipping: {sample_path} {uid} {person_id}")
# Mark for removal from Firestore (blob already gone)
samples_to_delete.append(sample_path)
continue
except Exception as e:
logger.error(f"Error downloading sample {sample_path}: {e} {uid} {person_id}")
# Transient download failure - keep sample, skip migration for now
has_transient_failures = True
continue
transcript, is_valid, reason = await verify_and_transcribe_sample(audio_bytes, 16000)
if is_valid:
valid_samples.append(sample_path)
valid_transcripts.append(transcript)
elif reason.startswith("transcription_failed"):
# Transient API failure - keep sample, don't migrate yet
logger.error(f"Transcription failed for {sample_path}, keeping sample: {reason} {uid} {person_id}")
has_transient_failures = True
else:
# Quality issue - mark for deletion (defer actual delete)
logger.info(f"Marking sample for deletion {sample_path}: {reason} {uid} {person_id}")
samples_to_delete.append(sample_path)
# Don't commit changes if there were transient failures - retry next time
if has_transient_failures:
logger.warning(f"Migration incomplete due to transient failures, will retry later {uid} {person_id}")
return person
# Now safe to delete blobs - no transient failures
for sample_path in samples_to_delete:
try:
await run_blocking(storage_executor, delete_sample_from_storage, sample_path)
except Exception as e:
logger.error(f"Failed to delete sample {sample_path}: {e} {uid} {person_id}")
new_embedding = None
if valid_samples:
try:
first_sample_audio = await run_blocking(storage_executor, download_sample_audio, valid_samples[0])
embedding = await run_blocking(
sync_executor, extract_embedding_from_bytes, first_sample_audio, "sample.wav"
)
new_embedding = embedding.flatten().tolist()
except Exception as e:
logger.error(f"Error extracting speaker embedding: {e} {uid} {person_id}")
await run_blocking(
db_executor,
users_db.update_person_speech_samples_after_migration,
uid,
person_id,
samples=valid_samples,
transcripts=valid_transcripts,
version=2,
speaker_embedding=new_embedding,
)
person['speech_samples'] = valid_samples
person['speech_sample_transcripts'] = valid_transcripts
person['speech_samples_version'] = 2
if new_embedding is not None:
person['speaker_embedding'] = new_embedding
elif not valid_samples:
person['speaker_embedding'] = None
return person
async def migrate_person_samples_v2_to_v3(uid: str, person: Dict[str, Any]) -> Dict[str, Any]:
"""
Migrate person's speech samples from v2 to v3.
v2: speech_samples + transcripts with v1 embeddings
v3: speech_samples + transcripts with v2 embeddings (regenerated)
Uses in-process lock to prevent concurrent migration for same person.
Args:
uid: User ID
person: Person dict with 'id', 'speech_samples', 'speech_samples_version', etc.
Returns:
Updated person dict with migrated fields
"""
version = person.get('speech_samples_version', 1)
if version >= 3:
return person
if version < 2:
# Need v1→v2 first
return person
person_id = person['id']
lock = await _get_migration_lock(uid, person_id)
async with lock:
# Re-check version inside lock (another call may have migrated)
fresh_person = await run_blocking(db_executor, users_db.get_person, uid, person_id)
if fresh_person and fresh_person.get('speech_samples_version', 1) >= 3:
return fresh_person
samples = person.get('speech_samples', [])
if not samples:
# No samples to re-extract from — clear stale embedding from old model
# first, then bump version (order matters: avoids race where a concurrent
# sample add writes a valid embedding that we'd then delete)
if person.get('speaker_embedding'):
await run_blocking(db_executor, users_db.clear_person_speaker_embedding, uid, person_id)
logger.info(f"v2→v3 migration: cleared stale embedding for person with no samples {uid} {person_id}")
await run_blocking(db_executor, users_db.update_person_speech_samples_version, uid, person_id, 3)
person['speech_samples_version'] = 3
person['speaker_embedding'] = None
return person
# Regenerate embedding from the first (latest) sample using v2/embedding API
new_embedding = None
try:
first_sample_audio = await run_blocking(storage_executor, download_sample_audio, samples[0])
embedding = await run_blocking(
sync_executor, extract_embedding_from_bytes, first_sample_audio, "sample.wav"
)
new_embedding = embedding.flatten().tolist()
except NotFound:
# Sample missing - don't advance to v3 to avoid caching stale v1 embedding
logger.warning(f"First sample not found during v2→v3 migration, skipping: {samples[0]} {uid} {person_id}")
return person
except Exception as e:
logger.error(f"Error extracting speaker embedding during v2→v3 migration: {e} {uid} {person_id}")
# Transient error, don't migrate yet
return person
# Update version and embedding
await run_blocking(
db_executor,
users_db.update_person_speech_samples_after_migration,
uid,
person_id,
samples=person.get('speech_samples', []),
transcripts=person.get('speech_sample_transcripts', []),
version=3,
speaker_embedding=new_embedding,
)
person['speech_samples_version'] = 3
person['speaker_embedding'] = new_embedding
return person
async def migrate_person_samples_v1_to_v3(uid: str, person: Dict[str, Any]) -> Dict[str, Any]:
"""
Migrate person's speech samples from v1 to v3.
This is a composite migration: v1 → v2 → v3.
v1: Only speech_samples (paths), no transcripts, v1 embeddings
v3: speech_samples + transcripts with v2 embeddings
Args:
uid: User ID
person: Person dict with 'id', 'speech_samples', 'speech_samples_version', etc.
Returns:
Updated person dict with migrated fields
"""
version = person.get('speech_samples_version', 1)
if version >= 3:
return person
# First do v1→v2 if needed
if version < 2:
person = await migrate_person_samples_v1_to_v2(uid, person)
# Check if v1→v2 succeeded
if person.get('speech_samples_version', 1) < 2:
return person # Transient failure, retry later
# Now do v2→v3
return await migrate_person_samples_v2_to_v3(uid, person)
async def maybe_migrate_person_samples(uid: str, person: Dict[str, Any]) -> Dict[str, Any]:
"""
Migrate person's speech samples to v3 if needed.
Checks speech_samples_version and triggers appropriate migration:
- v1 → v3 (composite through v2)
- v2 → v3
Args:
uid: User ID
person: Person dict
Returns:
Updated person dict (may be unchanged if already v3 or migration fails)
"""
version = person.get('speech_samples_version', 1)
if version >= 3:
return person
if version == 1:
return await migrate_person_samples_v1_to_v3(uid, person)
elif version == 2:
return await migrate_person_samples_v2_to_v3(uid, person)
return person