Forráskód Böngészése

feat(model): add models to stack gateways

Thomas Zhang 2 hónapja
szülő
commit
af217312a9

+ 90 - 0
internal/controller/stack_controller.go

@@ -28,11 +28,14 @@ import (
 	"k8s.io/apimachinery/pkg/api/resource"
 	metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
 	"k8s.io/apimachinery/pkg/runtime"
+	"k8s.io/apimachinery/pkg/types"
 	"k8s.io/utils/ptr"
 	ctrl "sigs.k8s.io/controller-runtime"
 	"sigs.k8s.io/controller-runtime/pkg/client"
 	"sigs.k8s.io/controller-runtime/pkg/controller/controllerutil"
+	"sigs.k8s.io/controller-runtime/pkg/handler"
 	logf "sigs.k8s.io/controller-runtime/pkg/log"
+	"sigs.k8s.io/controller-runtime/pkg/reconcile"
 
 	"github.com/LocoStack/loco-operator/api/v1alpha1"
 	"github.com/LocoStack/loco-operator/internal/reconciler"
@@ -71,6 +74,9 @@ func (r *StackReconciler) Reconcile(ctx context.Context, req ctrl.Request) (ctrl
 	if err := r.reconcileComponents(ctx, s); err != nil {
 		return ctrl.Result{}, err
 	}
+	if err := r.reconcileGateways(ctx, s); err != nil {
+		return ctrl.Result{}, err
+	}
 
 	s.Status.ObservedGeneration = s.Generation
 	if err := r.Status().Patch(ctx, s, patch); err != nil {
@@ -281,12 +287,96 @@ func (r *StackReconciler) reconcileComponent(ctx context.Context, s *v1alpha1.St
 	return tmpl.Name, available, err
 }
 
+func (r *StackReconciler) reconcileGateways(ctx context.Context, s *v1alpha1.Stack) error {
+	log := logf.FromContext(ctx)
+
+	gwList := &v1alpha1.ComponentList{}
+	err := r.List(ctx, gwList, client.InNamespace(s.Namespace), client.MatchingLabels{"locostack.com/stack": s.Name, "locostack.com/component": "Gateway"})
+	if err != nil {
+		return fmt.Errorf("failed to list gateway components: %v", err)
+	}
+	if len(gwList.Items) == 0 {
+		log.Info("No gateways found for stack", "namespace", s.Namespace, "stack", s.Name)
+		return nil
+	}
+
+	dependencies := make([]corev1.ObjectReference, 0)
+
+	emList := &v1alpha1.ExternalModelList{}
+	if err := r.List(ctx, emList, client.InNamespace(s.Namespace)); err != nil {
+		return fmt.Errorf("Failed to list external models: %v", err)
+	}
+	for _, em := range emList.Items {
+		if em.Spec.StackRef.Name == s.Name {
+			dependencies = append(dependencies, corev1.ObjectReference{
+				Kind:      "ExternalModel",
+				Name:      em.Name,
+				Namespace: em.Namespace,
+			})
+		}
+	}
+
+	mmList := &v1alpha1.ManagedModelList{}
+	if err := r.List(ctx, mmList, client.InNamespace(s.Namespace)); err != nil {
+		return fmt.Errorf("Failed to list managed models: %v", err)
+	}
+	for _, em := range mmList.Items {
+		if em.Spec.StackRef.Name == s.Name {
+			dependencies = append(dependencies, corev1.ObjectReference{
+				Kind:      "ManagedModel",
+				Name:      em.Name,
+				Namespace: em.Namespace,
+			})
+		}
+	}
+
+	if len(dependencies) == 0 {
+		return nil
+	}
+
+	for _, gw := range gwList.Items {
+		gw.Spec.Dependencies = dependencies
+		if err := r.Update(ctx, &gw); err != nil {
+			return fmt.Errorf("Failed to update gateway %s with dependencies: %v", gw.Name, err)
+		}
+	}
+
+	return nil
+}
+
 // SetupWithManager sets up the controller with the Manager.
 func (r *StackReconciler) SetupWithManager(mgr ctrl.Manager) error {
+	mapToStack := func(ctx context.Context, obj client.Object) []reconcile.Request {
+		var stackName string
+		switch typed := obj.(type) {
+		case *v1alpha1.ExternalModel:
+			stackName = typed.Spec.StackRef.Name
+		case *v1alpha1.ManagedModel:
+			stackName = typed.Spec.StackRef.Name
+		default:
+			return nil
+		}
+		stackList := &v1alpha1.StackList{}
+		if err := mgr.GetClient().List(ctx, stackList, client.InNamespace(obj.GetNamespace())); err != nil {
+			return nil
+		}
+		var reqs []reconcile.Request
+		for _, stack := range stackList.Items {
+			if stackName == stack.Name {
+				reqs = append(reqs, reconcile.Request{
+					NamespacedName: types.NamespacedName{Name: stack.Name, Namespace: stack.Namespace},
+				})
+			}
+		}
+		return reqs
+	}
+
 	return ctrl.NewControllerManagedBy(mgr).
 		For(&v1alpha1.Stack{}).
 		Owns(&v1alpha1.Component{}).
 		Owns(&corev1.PersistentVolumeClaim{}).
+		Watches(&v1alpha1.ExternalModel{}, handler.EnqueueRequestsFromMapFunc(mapToStack)).
+		Watches(&v1alpha1.ManagedModel{}, handler.EnqueueRequestsFromMapFunc(mapToStack)).
 		Named("stack").
 		Complete(r)
 }

+ 93 - 3
internal/reconciler/litellm.go

@@ -23,6 +23,7 @@ import (
 	"encoding/hex"
 	"fmt"
 	"maps"
+	"sort"
 
 	"github.com/LocoStack/loco-operator/api/v1alpha1"
 	"github.com/LocoStack/loco-operator/pkg/templates/litellm"
@@ -56,10 +57,41 @@ func (r *LiteLLMReconciler) ReconcileComponent(ctx context.Context, tmpl *v1alph
 	if err := r.ReconcileKey(ctx, "litellm", litellm.LITELLM_AUTH_SECRET_KEY, generateKey); err != nil {
 		return nil, fmt.Errorf("Failed to reconcile LiteLLM master key: %w", err)
 	}
-	configHash, err := r.reconcileConfigMap(ctx)
+	deps := make([]client.Object, 0)
+	emList := make([]*v1alpha1.ExternalModel, 0)
+	mmList := make([]*v1alpha1.ManagedModel, 0)
+	for _, dep := range r.component.Spec.Dependencies {
+		switch dep.Kind {
+		case "ExternalModel":
+			em := &v1alpha1.ExternalModel{}
+			if err := r.client.Get(ctx, client.ObjectKey{Name: dep.Name, Namespace: dep.Namespace}, em); err != nil {
+				return nil, fmt.Errorf("Failed to get ExternalModel %s/%s: %w", dep.Namespace, dep.Name, err)
+			}
+			emList = append(emList, em)
+			deps = append(deps, em)
+		case "ManagedModel":
+			mm := &v1alpha1.ManagedModel{}
+			if err := r.client.Get(ctx, client.ObjectKey{Name: dep.Name, Namespace: dep.Namespace}, mm); err != nil {
+				return nil, fmt.Errorf("Failed to get ManagedModel %s/%s: %w", dep.Namespace, dep.Name, err)
+			}
+			mmList = append(mmList, mm)
+			deps = append(deps, mm)
+		}
+	}
+	configHash, err := r.reconcileConfigMap(ctx, emList, mmList)
 	if err != nil {
 		return nil, fmt.Errorf("Failed to reconcile LiteLLM config: %w", err)
 	}
+	authSpecs := make(map[string]map[string]*v1alpha1.AuthSpec)
+	for _, em := range emList {
+		if em.Spec.Auth != nil {
+			if _, ok := authSpecs["ExternalModel"]; !ok {
+				authSpecs["ExternalModel"] = make(map[string]*v1alpha1.AuthSpec)
+			}
+			authSpecs["ExternalModel"][em.Name] = em.Spec.Auth
+		}
+	}
+	injectSecrets(tmpl, authSpecs)
 	if tmpl.Metadata.Annotations == nil {
 		tmpl.Metadata.Annotations = make(map[string]string)
 	}
@@ -69,10 +101,10 @@ func (r *LiteLLMReconciler) ReconcileComponent(ctx context.Context, tmpl *v1alph
 	if _, err := r.DefaultComponentReconciler.ReconcileComponent(ctx, tmpl, variables); err != nil {
 		return nil, err
 	}
-	return nil, nil
+	return deps, nil
 }
 
-func (r *LiteLLMReconciler) reconcileConfigMap(ctx context.Context) (string, error) {
+func (r *LiteLLMReconciler) reconcileConfigMap(ctx context.Context, emList []*v1alpha1.ExternalModel, mmList []*v1alpha1.ManagedModel) (string, error) {
 	var o11yComp *v1alpha1.Component
 	if r.stack.Spec.Observability != nil && r.stack.Spec.Observability.Enabled {
 		o11yComp = &v1alpha1.Component{}
@@ -84,6 +116,8 @@ func (r *LiteLLMReconciler) reconcileConfigMap(ctx context.Context) (string, err
 		MasterKeyEnvName:       litellm.LITELLM_MASTER_KEY_ENV_NAME,
 		Stack:                  r.stack,
 		ObservabilityComponent: o11yComp,
+		ExternalModels:         emList,
+		ManagedModels:          mmList,
 	}
 	config, err := configBuilder.BuildLiteLLMConfig()
 	if err != nil {
@@ -121,3 +155,59 @@ func generateKey() (string, error) {
 	}
 	return "sk-" + hex.EncodeToString(b), nil
 }
+
+func injectSecrets(tmpl *v1alpha1.Template, authSpecs map[string]map[string]*v1alpha1.AuthSpec) {
+	envVars := []corev1.EnvVar{}
+	kinds := make([]string, 0, len(authSpecs))
+	for kind := range authSpecs {
+		kinds = append(kinds, kind)
+	}
+	sort.Strings(kinds)
+	for _, kind := range kinds {
+		kindAuthSpecs := authSpecs[kind]
+		names := make([]string, 0, len(kindAuthSpecs))
+		for name := range kindAuthSpecs {
+			names = append(names, name)
+		}
+		sort.Strings(names)
+		for _, name := range names {
+			auth := kindAuthSpecs[name]
+			var refName string
+			var refKey string
+			if auth.APIKey != nil {
+				refName = auth.APIKey.SecretRef.Name
+				refKey = auth.APIKey.SecretRef.Key
+			} else if auth.BearerToken != nil {
+				refName = auth.BearerToken.Name
+				refKey = auth.BearerToken.Key
+			}
+			if refName != "" && refKey != "" {
+				envVars = append(envVars, corev1.EnvVar{
+					Name: litellm.AuthEnvVarName(kind, name),
+					ValueFrom: &corev1.EnvVarSource{
+						SecretKeyRef: &corev1.SecretKeySelector{
+							LocalObjectReference: corev1.LocalObjectReference{Name: refName},
+							Key:                  refKey,
+						},
+					},
+				})
+			}
+			if auth.Headers != nil {
+				for _, header := range auth.Headers {
+					if header.ValueFrom.SecretKeyRef != nil {
+						envVars = append(envVars, corev1.EnvVar{
+							Name: litellm.AuthEnvVarName(kind, name+"_"+header.Name),
+							ValueFrom: &corev1.EnvVarSource{
+								SecretKeyRef: &corev1.SecretKeySelector{
+									LocalObjectReference: header.ValueFrom.SecretKeyRef.LocalObjectReference,
+									Key:                  header.ValueFrom.SecretKeyRef.Key,
+								},
+							},
+						})
+					}
+				}
+			}
+		}
+	}
+	tmpl.Spec.Runtime.Env = append(tmpl.Spec.Runtime.Env, envVars...)
+}

+ 102 - 2
pkg/templates/litellm/config.go

@@ -16,7 +16,13 @@ limitations under the License.
 
 package litellm
 
-import "github.com/LocoStack/loco-operator/api/v1alpha1"
+import (
+	"fmt"
+	"strconv"
+	"strings"
+
+	"github.com/LocoStack/loco-operator/api/v1alpha1"
+)
 
 type LiteLLMConfig struct {
 	ModelList       []modelEntry              `yaml:"model_list"`
@@ -72,9 +78,23 @@ type LiteLLMConfigBuilder struct {
 	MasterKeyEnvName       string
 	Stack                  *v1alpha1.Stack
 	ObservabilityComponent *v1alpha1.Component
+	ExternalModels         []*v1alpha1.ExternalModel
+	ManagedModels          []*v1alpha1.ManagedModel
 }
 
 func (l *LiteLLMConfigBuilder) BuildLiteLLMConfig() (LiteLLMConfig, error) {
+	modelList := []modelEntry{}
+	for _, em := range l.ExternalModels {
+		entry := buildExternalModelEntry(em)
+		modelList = append(modelList, entry)
+	}
+	for _, mm := range l.ManagedModels {
+		if mm.Status.Endpoint == "" {
+			continue
+		}
+		entry := buildManagedModelEntry(mm)
+		modelList = append(modelList, entry)
+	}
 	phoenixEnabled := phoenixEnabled(l.Stack)
 	successCB := []string{}
 	failureCB := []string{}
@@ -83,7 +103,7 @@ func (l *LiteLLMConfigBuilder) BuildLiteLLMConfig() (LiteLLMConfig, error) {
 		failureCB = append(failureCB, "arize_phoenix")
 	}
 	cfg := LiteLLMConfig{
-		ModelList:  []modelEntry{},
+		ModelList:  modelList,
 		MCPServers: map[string]mcpServerEntry{},
 		LiteLLMSettings: &liteLLMSettings{
 			SuccessCallback: successCB,
@@ -118,3 +138,83 @@ func phoenixEnabled(stack *v1alpha1.Stack) bool {
 	}
 	return true
 }
+
+func AuthEnvVarName(kind string, name string) string {
+	return fmt.Sprintf("%s_%s_AUTH", strings.ToUpper(kind), strings.ToUpper(strings.ReplaceAll(name, "-", "_")))
+}
+
+// buildExternalModelEntry converts an ExternalModel into a LiteLLM model_list entry.
+func buildExternalModelEntry(em *v1alpha1.ExternalModel) modelEntry {
+	params := map[string]any{
+		"model": fmt.Sprintf("%s/%s", em.Spec.Provider, em.Spec.ProviderModel),
+	}
+	if em.Spec.Auth != nil && (em.Spec.Auth.APIKey != nil || em.Spec.Auth.BearerToken != nil) {
+		params["api_key"] = "os.environ/" + AuthEnvVarName("ExternalModel", em.Name)
+	}
+	if em.Spec.APIBase != "" {
+		params["api_base"] = em.Spec.APIBase
+	}
+	if em.Spec.APIVersion != "" {
+		params["api_version"] = em.Spec.APIVersion
+	}
+	switch em.Spec.Category {
+	case "embedding":
+		params["mode"] = "embedding"
+	case "reranker":
+		params["mode"] = "rerank"
+	}
+	for k, v := range em.Spec.ExtraParams {
+		params[k] = v
+	}
+	if em.Spec.DefaultInferenceParams != nil {
+		applyInferenceParams(params, em.Spec.DefaultInferenceParams)
+	}
+	return modelEntry{
+		ModelName:     em.Spec.ModelName,
+		LiteLLMParams: params,
+	}
+}
+
+// buildManagedModelEntry converts a ManagedModel into a LiteLLM model_list entry.
+func buildManagedModelEntry(mm *v1alpha1.ManagedModel) modelEntry {
+	// litellm's rerank dispatch does not recognize "openai" as a provider;
+	// "hosted_vllm" is the generic self-hosted provider for the /rerank call type.
+	provider := "openai"
+	if mm.Spec.Category == "reranker" {
+		provider = "hosted_vllm"
+	}
+	params := map[string]any{
+		"model":    fmt.Sprintf("%s/%s", provider, mm.Spec.ModelName),
+		"api_base": mm.Status.Endpoint + "/v1",
+		"api_key":  "not-set",
+	}
+	switch mm.Spec.Category {
+	case "embedding":
+		params["mode"] = "embedding"
+	case "reranker":
+		params["mode"] = "rerank"
+	}
+	return modelEntry{
+		ModelName:     mm.Spec.ModelName,
+		LiteLLMParams: params,
+	}
+}
+
+func applyInferenceParams(params map[string]any, inf *v1alpha1.InferenceParameters) {
+	if inf.Temperature != "" {
+		if v, err := strconv.ParseFloat(inf.Temperature, 64); err == nil {
+			params["temperature"] = v
+		}
+	}
+	if inf.TopP != "" {
+		if v, err := strconv.ParseFloat(inf.TopP, 64); err == nil {
+			params["top_p"] = v
+		}
+	}
+	if inf.MaxTokens != nil {
+		params["max_tokens"] = *inf.MaxTokens
+	}
+	for k, v := range inf.Extra {
+		params[k] = v
+	}
+}