@@ -23,6 +23,8 @@ import (
2323 "github.com/mudler/LocalAI/core/schema"
2424 "github.com/mudler/LocalAI/core/services/galleryop"
2525 "github.com/mudler/LocalAI/core/services/modeladmin"
26+ "github.com/mudler/LocalAI/core/services/nodes"
27+ "github.com/mudler/LocalAI/core/services/nodes/prefixcache"
2628 "github.com/mudler/LocalAI/core/services/routing/billing"
2729 "github.com/mudler/LocalAI/core/services/routing/pii"
2830 "github.com/mudler/LocalAI/core/services/routing/router"
@@ -47,6 +49,7 @@ type Client struct {
4749 ConfigLoader * config.ModelConfigLoader
4850 ModelLoader * model.ModelLoader
4951 Gallery * galleryop.GalleryService
52+ NodeRegistry * nodes.NodeRegistry
5053 VoiceProfiles * voiceprofile.Store
5154
5255 // StatsRecorder and FallbackUser are optional — they back the
@@ -75,13 +78,18 @@ type Client struct {
7578// except ModelLoader (used only for SystemInfo's loaded-models report and
7679// best-effort ShutdownModel calls during config edits) and the stats
7780// fields (StatsRecorder, FallbackUser) which gate get_usage_stats.
78- func New (appConfig * config.ApplicationConfig , systemState * system.SystemState , cl * config.ModelConfigLoader , ml * model.ModelLoader , gs * galleryop.GalleryService ) * Client {
81+ func New (appConfig * config.ApplicationConfig , systemState * system.SystemState , cl * config.ModelConfigLoader , ml * model.ModelLoader , gs * galleryop.GalleryService , registries ... * nodes.NodeRegistry ) * Client {
82+ var registry * nodes.NodeRegistry
83+ if len (registries ) > 0 {
84+ registry = registries [0 ]
85+ }
7986 return & Client {
8087 AppConfig : appConfig ,
8188 SystemState : systemState ,
8289 ConfigLoader : cl ,
8390 ModelLoader : ml ,
8491 Gallery : gs ,
92+ NodeRegistry : registry ,
8593 VoiceProfiles : voiceprofile .NewStore (appConfig .DataPath ),
8694 modelAdmin : modeladmin .NewConfigService (cl , appConfig ),
8795 }
@@ -501,20 +509,119 @@ func (c *Client) ListNodes(_ context.Context) ([]localaitools.Node, error) {
501509 return []localaitools.Node {}, nil
502510}
503511
504- func (c * Client ) ListScheduling (_ context.Context ) ([]localaitools.ModelSchedulingConfig , error ) {
505- return []localaitools.ModelSchedulingConfig {}, nil
512+ func (c * Client ) ListScheduling (ctx context.Context ) ([]localaitools.ModelSchedulingConfig , error ) {
513+ if c .NodeRegistry == nil {
514+ return []localaitools.ModelSchedulingConfig {}, nil
515+ }
516+ configs , err := c .NodeRegistry .ListModelSchedulings (ctx )
517+ if err != nil {
518+ return nil , err
519+ }
520+ return localaitools .SchedulingConfigsFromNodes (configs ), nil
506521}
507522
508- func (c * Client ) GetScheduling (_ context.Context , _ string ) (* localaitools.ModelSchedulingConfig , error ) {
509- return nil , errors .New ("model scheduling is only available in distributed mode" )
523+ func (c * Client ) GetScheduling (ctx context.Context , modelName string ) (* localaitools.ModelSchedulingConfig , error ) {
524+ if modelName == "" {
525+ return nil , errors .New ("model_name is required" )
526+ }
527+ if c .NodeRegistry == nil {
528+ return nil , errors .New ("model scheduling is only available in distributed mode" )
529+ }
530+ config , err := c .NodeRegistry .GetModelScheduling (ctx , modelName )
531+ if err != nil {
532+ return nil , err
533+ }
534+ if config == nil {
535+ return nil , nil
536+ }
537+ out := localaitools .SchedulingConfigFromNode (* config )
538+ return & out , nil
510539}
511540
512- func (c * Client ) SetScheduling (_ context.Context , _ localaitools.SetSchedulingRequest ) (* localaitools.ModelSchedulingConfig , error ) {
513- return nil , errors .New ("model scheduling is only available in distributed mode" )
541+ func (c * Client ) SetScheduling (ctx context.Context , req localaitools.SetSchedulingRequest ) (* localaitools.ModelSchedulingConfig , error ) {
542+ if req .ModelName == "" {
543+ return nil , errors .New ("model_name is required" )
544+ }
545+ if c .NodeRegistry == nil {
546+ return nil , errors .New ("model scheduling is only available in distributed mode" )
547+ }
548+
549+ existing , err := c .NodeRegistry .GetModelScheduling (ctx , req .ModelName )
550+ if err != nil {
551+ return nil , fmt .Errorf ("load existing scheduling config: %w" , err )
552+ }
553+ routePolicy := ""
554+ absThr := 0
555+ relThr := 0.0
556+ minMatch := 0.0
557+ if existing != nil {
558+ routePolicy = existing .RoutePolicy
559+ absThr = existing .BalanceAbsThreshold
560+ relThr = existing .BalanceRelThreshold
561+ minMatch = existing .MinPrefixMatch
562+ }
563+ if req .RoutePolicy != nil {
564+ routePolicy = * req .RoutePolicy
565+ }
566+ if req .BalanceAbsThreshold != nil {
567+ absThr = * req .BalanceAbsThreshold
568+ }
569+ if req .BalanceRelThreshold != nil {
570+ relThr = * req .BalanceRelThreshold
571+ }
572+ if req .MinPrefixMatch != nil {
573+ minMatch = * req .MinPrefixMatch
574+ }
575+ if req .SpreadAll && (req .MinReplicas != 0 || req .MaxReplicas != 0 ) {
576+ return nil , errors .New ("spread_all and min_replicas/max_replicas are mutually exclusive" )
577+ }
578+ if req .MinReplicas < 0 {
579+ return nil , errors .New ("min_replicas must be >= 0" )
580+ }
581+ if req .MaxReplicas < 0 {
582+ return nil , errors .New ("max_replicas must be >= 0" )
583+ }
584+ if req .MaxReplicas > 0 && req .MinReplicas > req .MaxReplicas {
585+ return nil , errors .New ("min_replicas must be <= max_replicas" )
586+ }
587+ if err := prefixcache .ValidateThresholds (routePolicy , absThr , relThr , minMatch ); err != nil {
588+ return nil , err
589+ }
590+
591+ var selectorJSON string
592+ if len (req .NodeSelector ) > 0 {
593+ b , err := json .Marshal (req .NodeSelector )
594+ if err != nil {
595+ return nil , fmt .Errorf ("invalid node_selector: %w" , err )
596+ }
597+ selectorJSON = string (b )
598+ }
599+ config := & nodes.ModelSchedulingConfig {
600+ ModelName : req .ModelName ,
601+ NodeSelector : selectorJSON ,
602+ MinReplicas : req .MinReplicas ,
603+ MaxReplicas : req .MaxReplicas ,
604+ SpreadAll : req .SpreadAll ,
605+ RoutePolicy : routePolicy ,
606+ BalanceAbsThreshold : absThr ,
607+ BalanceRelThreshold : relThr ,
608+ MinPrefixMatch : minMatch ,
609+ }
610+ if err := c .NodeRegistry .SetModelScheduling (ctx , config ); err != nil {
611+ return nil , err
612+ }
613+ out := localaitools .SchedulingConfigFromNode (* config )
614+ return & out , nil
514615}
515616
516- func (c * Client ) DeleteScheduling (_ context.Context , _ string ) error {
517- return errors .New ("model scheduling is only available in distributed mode" )
617+ func (c * Client ) DeleteScheduling (ctx context.Context , modelName string ) error {
618+ if modelName == "" {
619+ return errors .New ("model_name is required" )
620+ }
621+ if c .NodeRegistry == nil {
622+ return errors .New ("model scheduling is only available in distributed mode" )
623+ }
624+ return c .NodeRegistry .DeleteModelScheduling (ctx , modelName )
518625}
519626
520627func (c * Client ) SetNodeVRAMBudget (_ context.Context , _ , _ string ) error {
0 commit comments