summaryrefslogtreecommitdiff
path: root/internal/catalog
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/catalog/catalog.go19
1 files changed, 19 insertions, 0 deletions
diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go
index 362d5b0..32024aa 100644
--- a/internal/catalog/catalog.go
+++ b/internal/catalog/catalog.go
@@ -42,6 +42,7 @@ func New(cfg config.Config) *Catalog {
model := domain.Model{
ID: modelCfg.ID,
OwnedBy: modelCfg.OwnedBy,
+ Capabilities: append([]string(nil), modelCfg.Capabilities...),
InputPriceMicrosPerMillion: modelCfg.InputPriceMicrosPerMillion,
OutputPriceMicrosPerMillion: modelCfg.OutputPriceMicrosPerMillion,
CacheReadPriceMicrosPerMillion: modelCfg.CacheReadPriceMicrosPerMillion,
@@ -139,12 +140,30 @@ func (c *Catalog) Models(protocol domain.Protocol) []domain.Model {
return result
}
+// AllModels returns a detached view of the current atomic catalog snapshot.
+// It is intended for control-loop work such as active provider probes; request
+// routing should continue to use ModelForPrincipal and ModelsFor.
+func (c *Catalog) AllModels() []domain.Model {
+ current := c.state.Load()
+ if current == nil {
+ return nil
+ }
+ result := make([]domain.Model, len(current.list))
+ for index, model := range current.list {
+ result[index] = model
+ result[index].Routes = append([]domain.Route(nil), model.Routes...)
+ }
+ return result
+}
+
func protocolCompatible(provider domain.Provider, requestProtocol domain.Protocol) bool {
switch requestProtocol {
case domain.ProtocolOpenAI:
return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions"
case domain.ProtocolOpenAIResponses:
return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses"
+ case domain.ProtocolOpenAIEmbeddings:
+ return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "embeddings"
case domain.ProtocolAnthropic:
return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages"
default: