@@ -16,54 +16,94 @@ def _clean_text(value: Any) -> Optional[str]:
1616 return text or None
1717
1818
19- def _resolve_voice_ref (model_root : str , value : Any ) -> Optional [str ]:
19+ def _reference_roots (model_root : str , reference_root : Optional [str ] = None ) -> list [str ]:
20+ roots : list [str ] = []
21+ for root in (reference_root , model_root ):
22+ text = _clean_text (root )
23+ if text :
24+ resolved = os .path .abspath (text )
25+ if resolved not in roots :
26+ roots .append (resolved )
27+ return roots
28+
29+
30+ def _resolve_voice_ref (
31+ model_root : str ,
32+ value : Any ,
33+ * ,
34+ reference_root : Optional [str ] = None ,
35+ ) -> Optional [str ]:
2036 raw = _clean_text (value )
2137 if not raw :
2238 return None
2339 if os .path .isabs (raw ):
2440 return raw
25- root = os .path .abspath (model_root )
41+ roots = _reference_roots (model_root , reference_root )
42+ for root in roots :
43+ candidate = os .path .abspath (os .path .join (root , raw ))
44+ if os .path .exists (candidate ):
45+ return candidate
46+ root = roots [0 ] if roots else os .path .abspath (model_root )
2647 return os .path .abspath (os .path .join (root , raw ))
2748
2849
29- def _relative_voice_ref_escapes (model_root : str , value : Any ) -> bool :
50+ def _relative_voice_ref_escapes (
51+ model_root : str ,
52+ value : Any ,
53+ * ,
54+ reference_root : Optional [str ] = None ,
55+ ) -> bool :
3056 raw = _clean_text (value )
3157 if not raw or os .path .isabs (raw ):
3258 return False
33- root = os .path .realpath (model_root )
34- target = os .path .realpath (os .path .join (root , raw ))
35- try :
36- return os .path .commonpath ([root , target ]) != root
37- except ValueError :
38- return True
59+ escaped_all_roots = True
60+ for root in _reference_roots (model_root , reference_root ):
61+ root_real = os .path .realpath (root )
62+ target = os .path .realpath (os .path .join (root_real , raw ))
63+ try :
64+ if os .path .commonpath ([root_real , target ]) == root_real :
65+ escaped_all_roots = False
66+ except ValueError :
67+ continue
68+ return escaped_all_roots
3969
4070
4171def _validate_voice_ref_path (
4272 * ,
4373 label : str ,
4474 value : Any ,
4575 model_root : str ,
76+ reference_root : Optional [str ] = None ,
4677 errors : list [str ],
4778) -> None :
4879 raw = _clean_text (value )
4980 if not raw :
5081 return
51- if _relative_voice_ref_escapes (model_root , raw ):
82+ if _relative_voice_ref_escapes (model_root , raw , reference_root = reference_root ):
5283 errors .append (f"{ label } escapes model bundle: { raw } " )
5384 return
54- resolved = _resolve_voice_ref (model_root , raw )
85+ resolved = _resolve_voice_ref (model_root , raw , reference_root = reference_root )
5586 if not resolved or not os .path .exists (resolved ):
5687 errors .append (f"{ label } does not exist: { resolved or raw } " )
5788
5889
59- def normalize_voice_preset (preset : Any , * , model_root : str ) -> Optional [Dict [str , str ]]:
90+ def normalize_voice_preset (
91+ preset : Any ,
92+ * ,
93+ model_root : str ,
94+ reference_root : Optional [str ] = None ,
95+ ) -> Optional [Dict [str , str ]]:
6096 if not isinstance (preset , dict ):
6197 return None
6298 out : Dict [str , str ] = {}
6399 voice_id = _clean_text (preset .get ("voice_id" ))
64100 if voice_id :
65101 out ["voice_id" ] = voice_id
66- voice_ref = _resolve_voice_ref (model_root , preset .get ("voice_ref" ))
102+ voice_ref = _resolve_voice_ref (
103+ model_root ,
104+ preset .get ("voice_ref" ),
105+ reference_root = reference_root ,
106+ )
67107 if voice_ref :
68108 out ["voice_ref" ] = voice_ref
69109 reference_text = _clean_text (preset .get ("reference_text" ))
@@ -76,6 +116,7 @@ def normalize_voice_presets(
76116 presets : Any ,
77117 * ,
78118 model_root : str ,
119+ reference_root : Optional [str ] = None ,
79120) -> Dict [str , Dict [str , str ]]:
80121 if not isinstance (presets , dict ):
81122 return {}
@@ -84,7 +125,11 @@ def normalize_voice_presets(
84125 name = str (raw_name or "" ).strip ()
85126 if not name :
86127 continue
87- normalized = normalize_voice_preset (raw_preset , model_root = model_root )
128+ normalized = normalize_voice_preset (
129+ raw_preset ,
130+ model_root = model_root ,
131+ reference_root = reference_root ,
132+ )
88133 if normalized :
89134 out [name ] = normalized
90135 return out
@@ -94,6 +139,7 @@ def normalize_default_voice_preset(
94139 value : Any ,
95140 * ,
96141 model_root : str ,
142+ reference_root : Optional [str ] = None ,
97143 voice_presets : Optional [Dict [str , Dict [str , str ]]] = None ,
98144) -> Any :
99145 if value is None or value == "" :
@@ -102,7 +148,11 @@ def normalize_default_voice_preset(
102148 name = value .strip ()
103149 return name or None
104150 if isinstance (value , dict ):
105- normalized = normalize_voice_preset (value , model_root = model_root )
151+ normalized = normalize_voice_preset (
152+ value ,
153+ model_root = model_root ,
154+ reference_root = reference_root ,
155+ )
106156 return normalized
107157 return None
108158
@@ -173,7 +223,30 @@ def _task_normalized_to_swap_params(normalized: Dict[str, Any]) -> Dict[str, Any
173223 return params
174224
175225
176- def audio_request_defaults_to_swap_set_params (config : dict ) -> Dict [str , Any ]:
226+ def _resolve_normalized_reference_fields (
227+ normalized : Dict [str , Any ],
228+ * ,
229+ model_root : Optional [str ] = None ,
230+ reference_root : Optional [str ] = None ,
231+ ) -> Dict [str , Any ]:
232+ if not model_root or "voice_ref" not in normalized :
233+ return normalized
234+ resolved = _resolve_voice_ref (
235+ model_root ,
236+ normalized .get ("voice_ref" ),
237+ reference_root = reference_root ,
238+ )
239+ if not resolved :
240+ return normalized
241+ return {** normalized , "voice_ref" : resolved }
242+
243+
244+ def audio_request_defaults_to_swap_set_params (
245+ config : dict ,
246+ * ,
247+ model_root : Optional [str ] = None ,
248+ reference_root : Optional [str ] = None ,
249+ ) -> Dict [str , Any ]:
177250 """
178251 Map Studio request defaults to llama-swap ``filters.setParams`` for audio.cpp.
179252
@@ -186,6 +259,11 @@ def audio_request_defaults_to_swap_set_params(config: dict) -> Dict[str, Any]:
186259
187260 defaults_key = request_defaults_key_for (config .get ("task" ), config .get ("family" ))
188261 normalized = normalize_request_defaults (defaults_key , config .get (defaults_key ))
262+ normalized = _resolve_normalized_reference_fields (
263+ normalized ,
264+ model_root = model_root ,
265+ reference_root = reference_root ,
266+ )
189267 if not normalized :
190268 return {}
191269 if defaults_key == "speech_defaults" :
@@ -254,16 +332,25 @@ def validate_voice_presets(
254332 config : dict ,
255333 * ,
256334 model_root : str ,
335+ reference_root : Optional [str ] = None ,
257336 errors : list [str ],
258337) -> Dict [str , Dict [str , str ]]:
259- presets = normalize_voice_presets (config .get ("voice_presets" ), model_root = model_root )
338+ presets = normalize_voice_presets (
339+ config .get ("voice_presets" ),
340+ model_root = model_root ,
341+ reference_root = reference_root ,
342+ )
260343 default_value = config .get ("default_voice_preset" )
261344 if isinstance (default_value , str ):
262345 name = default_value .strip ()
263346 if name and name not in presets :
264347 errors .append (f"default_voice_preset '{ name } ' is not defined in voice_presets" )
265348 elif isinstance (default_value , dict ):
266- if not normalize_voice_preset (default_value , model_root = model_root ):
349+ if not normalize_voice_preset (
350+ default_value ,
351+ model_root = model_root ,
352+ reference_root = reference_root ,
353+ ):
267354 errors .append (
268355 "default_voice_preset must include voice_id, voice_ref, or reference_text"
269356 )
@@ -273,6 +360,7 @@ def validate_voice_presets(
273360 label = f"default_voice_preset { path_key } " ,
274361 value = raw ,
275362 model_root = model_root ,
363+ reference_root = reference_root ,
276364 errors = errors ,
277365 )
278366 raw_presets = config .get ("voice_presets" )
@@ -288,6 +376,7 @@ def validate_voice_presets(
288376 label = f"voice_presets.{ name } .voice_ref" ,
289377 value = raw_voice_ref ,
290378 model_root = model_root ,
379+ reference_root = reference_root ,
291380 errors = errors ,
292381 )
293382 return presets
@@ -297,6 +386,7 @@ def validate_speech_default_references(
297386 config : dict ,
298387 * ,
299388 model_root : str ,
389+ reference_root : Optional [str ] = None ,
300390 errors : list [str ],
301391) -> None :
302392 defaults = config .get ("speech_defaults" )
@@ -306,5 +396,6 @@ def validate_speech_default_references(
306396 label = "speech_defaults.voice_ref" ,
307397 value = defaults .get ("voice_ref" ),
308398 model_root = model_root ,
399+ reference_root = reference_root ,
309400 errors = errors ,
310401 )
0 commit comments