@@ -35,14 +35,14 @@ fn normalize_openai_base(override_base: &str) -> String {
3535 }
3636}
3737
38- fn canonical_local_provider_id ( provider_id : & str ) -> & str {
39- match provider_id {
40- "lmstudio" => ProviderId :: LM_STUDIO ,
41- "llamacpp" | "llama-server" => ProviderId :: LLAMA_CPP ,
42- _ => provider_id,
43- }
44- }
45-
38+ fn canonical_local_provider_id ( provider_id : & str ) -> & str {
39+ match provider_id {
40+ "lmstudio" => ProviderId :: LM_STUDIO ,
41+ "llamacpp" | "llama-server" => ProviderId :: LLAMA_CPP ,
42+ _ => provider_id,
43+ }
44+ }
45+
4646pub fn resolve_provider_api_base (
4747 config : & claurst_core:: config:: Config ,
4848 provider_id : & str ,
@@ -321,26 +321,26 @@ impl ProviderRegistry {
321321 ///
322322 /// # Panics
323323 /// Panics if no provider with that ID has been registered.
324- pub fn set_default ( & mut self , id : ProviderId ) -> & mut Self {
325- let canonical_id = ProviderId :: new ( canonical_local_provider_id ( & id) ) ;
326- assert ! (
327- self . providers. contains_key( & canonical_id) ,
328- "set_default: provider '{}' is not registered" ,
329- id,
330- ) ;
331- self . default_provider_id = canonical_id;
332- self
333- }
334-
335- /// Get a provider by ID.
336- pub fn get ( & self , id : & ProviderId ) -> Option < & Arc < dyn LlmProvider > > {
337- self . providers . get ( id) . or_else ( || {
338- let canonical_id = canonical_local_provider_id ( id) ;
339- ( canonical_id != & * * id)
340- . then ( || self . providers . get ( & ProviderId :: new ( canonical_id) ) )
341- . flatten ( )
342- } )
343- }
324+ pub fn set_default ( & mut self , id : ProviderId ) -> & mut Self {
325+ let canonical_id = ProviderId :: new ( canonical_local_provider_id ( & id) ) ;
326+ assert ! (
327+ self . providers. contains_key( & canonical_id) ,
328+ "set_default: provider '{}' is not registered" ,
329+ id,
330+ ) ;
331+ self . default_provider_id = canonical_id;
332+ self
333+ }
334+
335+ /// Get a provider by ID.
336+ pub fn get ( & self , id : & ProviderId ) -> Option < & Arc < dyn LlmProvider > > {
337+ self . providers . get ( id) . or_else ( || {
338+ let canonical_id = canonical_local_provider_id ( id) ;
339+ ( canonical_id != & * * id)
340+ . then ( || self . providers . get ( & ProviderId :: new ( canonical_id) ) )
341+ . flatten ( )
342+ } )
343+ }
344344
345345 /// Get the default provider.
346346 pub fn default_provider ( & self ) -> Option < & Arc < dyn LlmProvider > > {
@@ -665,35 +665,35 @@ impl Default for ProviderRegistry {
665665 Self :: new ( )
666666 }
667667}
668-
669- #[ cfg( test) ]
670- mod tests {
671- use super :: * ;
672- use crate :: providers;
673-
674- #[ test]
675- fn local_provider_aliases_resolve_to_canonical_registrations ( ) {
676- let mut registry = ProviderRegistry :: new ( ) ;
677- registry. register ( Arc :: new ( providers:: lm_studio ( ) ) ) ;
678- registry. register ( Arc :: new ( providers:: llama_cpp ( ) ) ) ;
679-
680- let lm_studio = registry
681- . get ( & ProviderId :: new ( "lmstudio" ) )
682- . expect ( "lmstudio alias should resolve" ) ;
683- let llama_cpp = registry
684- . get ( & ProviderId :: new ( "llamacpp" ) )
685- . expect ( "llamacpp alias should resolve" ) ;
686-
687- assert_eq ! ( & * * lm_studio. id( ) , ProviderId :: LM_STUDIO ) ;
688- assert_eq ! ( & * * llama_cpp. id( ) , ProviderId :: LLAMA_CPP ) ;
689- }
690-
691- #[ test]
692- fn alias_can_select_canonical_default_provider ( ) {
693- let mut registry = ProviderRegistry :: new ( ) ;
694- registry. register ( Arc :: new ( providers:: lm_studio ( ) ) ) ;
695- registry. set_default ( ProviderId :: new ( "lmstudio" ) ) ;
696-
697- assert_eq ! ( & * * registry. default_provider_id( ) , ProviderId :: LM_STUDIO ) ;
698- }
699- }
668+
669+ #[ cfg( test) ]
670+ mod tests {
671+ use super :: * ;
672+ use crate :: providers;
673+
674+ #[ test]
675+ fn local_provider_aliases_resolve_to_canonical_registrations ( ) {
676+ let mut registry = ProviderRegistry :: new ( ) ;
677+ registry. register ( Arc :: new ( providers:: lm_studio ( ) ) ) ;
678+ registry. register ( Arc :: new ( providers:: llama_cpp ( ) ) ) ;
679+
680+ let lm_studio = registry
681+ . get ( & ProviderId :: new ( "lmstudio" ) )
682+ . expect ( "lmstudio alias should resolve" ) ;
683+ let llama_cpp = registry
684+ . get ( & ProviderId :: new ( "llamacpp" ) )
685+ . expect ( "llamacpp alias should resolve" ) ;
686+
687+ assert_eq ! ( & * * lm_studio. id( ) , ProviderId :: LM_STUDIO ) ;
688+ assert_eq ! ( & * * llama_cpp. id( ) , ProviderId :: LLAMA_CPP ) ;
689+ }
690+
691+ #[ test]
692+ fn alias_can_select_canonical_default_provider ( ) {
693+ let mut registry = ProviderRegistry :: new ( ) ;
694+ registry. register ( Arc :: new ( providers:: lm_studio ( ) ) ) ;
695+ registry. set_default ( ProviderId :: new ( "lmstudio" ) ) ;
696+
697+ assert_eq ! ( & * * registry. default_provider_id( ) , ProviderId :: LM_STUDIO ) ;
698+ }
699+ }
0 commit comments