Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 62 additions & 11 deletions cmd/late/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -216,16 +216,23 @@ func main() {
Model: resolvedOpenAIConfig.Model,
EnableImages: *enableImagesReq,
}
if appConfig != nil {
if setting, ok := appConfig.GetModelForAgent("orchestrator"); ok {
resolvedClientConfig.BaseURL = setting.URL
resolvedClientConfig.APIKey = setting.Key
resolvedClientConfig.Model = setting.Model
}
}
c := client.NewClient(resolvedClientConfig)
c.DiscoverBackend(context.Background())

// Initialize Subagent Client
resolvedSubagentConfig := appconfig.ResolveSubagentSettings(appConfig, resolvedOpenAIConfig)

subagentClient := c
if resolvedSubagentConfig.BaseURL != resolvedOpenAIConfig.BaseURL ||
resolvedSubagentConfig.APIKey != resolvedOpenAIConfig.APIKey ||
resolvedSubagentConfig.Model != resolvedOpenAIConfig.Model {
if resolvedSubagentConfig.BaseURL != resolvedClientConfig.BaseURL ||
resolvedSubagentConfig.APIKey != resolvedClientConfig.APIKey ||
resolvedSubagentConfig.Model != resolvedClientConfig.Model {
subagentClient = client.NewClient(client.Config{
BaseURL: resolvedSubagentConfig.BaseURL,
APIKey: resolvedSubagentConfig.APIKey,
Expand Down Expand Up @@ -286,16 +293,33 @@ func main() {
// We'll add middlewares later once the program is started
rootAgent := orchestrator.NewBaseOrchestrator("main", sess, nil, 0)

model := tui.NewModel(rootAgent, renderer)
model.ModelName = resolvedOpenAIConfig.Model
model.ShowCWD = *showCWDReq
model := tui.NewModel(rootAgent, renderer, appConfig)
model.ApplyOrchestratorModel = func(setting appconfig.ModelSetting) {
sess.SetClient(newModelClient(context.Background(), setting, *enableImagesReq))
}
if appConfig != nil {
if orchestratorModel, ok := appConfig.AgentModels["orchestrator"]; ok {
model.ModelName = orchestratorModel
} else {
model.ModelName = resolvedOpenAIConfig.Model
}

// Detect if subagents use a different model/backend
if resolvedSubagentConfig.BaseURL != resolvedOpenAIConfig.BaseURL ||
resolvedSubagentConfig.APIKey != resolvedOpenAIConfig.APIKey ||
resolvedSubagentConfig.Model != resolvedOpenAIConfig.Model {
var subagentInfos []string
for _, sub := range assets.GetSubagents() {
if m, ok := appConfig.AgentModels[sub.Name]; ok {
subagentInfos = append(subagentInfos, fmt.Sprintf("%s:%s", sub.Name, m))
}
}
if len(subagentInfos) > 0 {
model.SubagentInfo = strings.Join(subagentInfos, ", ")
} else {
model.SubagentInfo = resolvedSubagentConfig.Model
}
} else {
model.ModelName = resolvedOpenAIConfig.Model
model.SubagentInfo = resolvedSubagentConfig.Model
}
model.ShowCWD = *showCWDReq

p := tea.NewProgram(model)

Expand All @@ -322,7 +346,23 @@ func main() {

if *enableSubagentsReq {
runner := func(ctx context.Context, goal string, ctxFiles []string, agentType string) (string, error) {
child, err := agent.NewSubagentOrchestrator(subagentClient, goal, ctxFiles, agentType, enabledTools, *injectCWDReq, *gemmaThinkingReq, *subagentMaxTurns, rootAgent, p)
var currentSubagentClient *client.Client
if appConfig != nil {
if setting, ok := appConfig.GetModelForAgent(agentType); ok {
currentSubagentClient = client.NewClient(client.Config{
BaseURL: setting.URL,
APIKey: setting.Key,
Model: setting.Model,
EnableImages: *enableImagesReq,
})
currentSubagentClient.DiscoverBackend(ctx)
}
}
if currentSubagentClient == nil {
currentSubagentClient = subagentClient
}

child, err := agent.NewSubagentOrchestrator(currentSubagentClient, goal, ctxFiles, agentType, enabledTools, *injectCWDReq, *gemmaThinkingReq, *subagentMaxTurns, rootAgent, p)
if err != nil {
return "", err
}
Expand Down Expand Up @@ -350,6 +390,17 @@ func main() {
}
}

func newModelClient(ctx context.Context, setting appconfig.ModelSetting, enableImages bool) *client.Client {
c := client.NewClient(client.Config{
BaseURL: setting.URL,
APIKey: setting.Key,
Model: setting.Model,
EnableImages: enableImages,
})
c.DiscoverBackend(ctx)
return c
}

// handleSessionCommand processes session subcommands
// Returns: command, args (remaining), verbose flag
func handleSessionCommand(args []string) (string, []string, bool) {
Expand Down
44 changes: 44 additions & 0 deletions internal/config/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,12 @@ type SubagentSettings struct {
Model string
}

type ModelSetting struct {
URL string `json:"url"`
Key string `json:"key"`
Model string `json:"model"`
}

const (
configDirPerm os.FileMode = 0o700
configFilePerm os.FileMode = 0o600
Expand All @@ -46,6 +52,9 @@ type Config struct {
SubagentModel string `json:"subagent_model,omitempty"`

SkillsDir string `json:"skills_dir,omitempty"`

Models []ModelSetting `json:"models,omitempty"`
AgentModels map[string]string `json:"agent_models,omitempty"`
}

func defaultConfig() Config {
Expand Down Expand Up @@ -240,3 +249,38 @@ func tightenPermission(path string, required os.FileMode) error {

return os.Chmod(path, required)
}

// GetModelForAgent returns the ModelSetting for a given agent type.
// If not found, it returns false.
func (cfg *Config) GetModelForAgent(agentType string) (ModelSetting, bool) {
if cfg == nil || cfg.AgentModels == nil || cfg.Models == nil {
return ModelSetting{}, false
}
modelName, exists := cfg.AgentModels[agentType]
if !exists {
return ModelSetting{}, false
}
for _, m := range cfg.Models {
if m.Model == modelName {
return m, true
}
}
return ModelSetting{}, false
}

// SaveConfig writes the configuration back to config.json.
func SaveConfig(cfg *Config) error {
lateConfigDir, err := pathutil.LateConfigDir()
if err != nil {
return err
}
configPath := filepath.Join(lateConfigDir, "config.json")
data, err := json.MarshalIndent(cfg, "", " ")
if err != nil {
return err
}
if err := os.WriteFile(configPath, data, configFilePerm); err != nil {
return err
}
return ensureSecureConfigPermissions(lateConfigDir, configPath)
}
37 changes: 37 additions & 0 deletions internal/config/config_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -527,3 +527,40 @@ func TestResolveSubagentSettings(t *testing.T) {
})
}
}

func TestConfig_GetModelForAgent(t *testing.T) {
cfg := &Config{
Models: []ModelSetting{
{URL: "http://localhost:8080", Key: "key-1", Model: "model-1"},
{URL: "http://localhost:9090", Key: "key-2", Model: "model-2"},
},
AgentModels: map[string]string{
"orchestrator": "model-1",
"coder": "model-2",
"unknown": "model-3",
},
}

tests := []struct {
agentType string
wantModel string
wantOk bool
}{
{"orchestrator", "model-1", true},
{"coder", "model-2", true},
{"unknown", "", false},
{"missing", "", false},
}

for _, tt := range tests {
t.Run(tt.agentType, func(t *testing.T) {
got, ok := cfg.GetModelForAgent(tt.agentType)
if ok != tt.wantOk {
t.Errorf("GetModelForAgent(%q) ok = %v, want %v", tt.agentType, ok, tt.wantOk)
}
if ok && got.Model != tt.wantModel {
t.Errorf("GetModelForAgent(%q) got model = %q, want %q", tt.agentType, got.Model, tt.wantModel)
}
})
}
}
18 changes: 18 additions & 0 deletions internal/session/client_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
package session

import (
"late/internal/client"
"testing"
)

func TestSetClient(t *testing.T) {
original := client.NewClient(client.Config{Model: "original"})
replacement := client.NewClient(client.Config{Model: "replacement"})
s := New(original, "", nil, "", false)

s.SetClient(replacement)

if got := s.Client(); got != replacement {
t.Fatalf("Client() = %p, want replacement %p", got, replacement)
}
}
18 changes: 15 additions & 3 deletions internal/session/session.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,11 +9,13 @@ import (
"late/internal/tool"
"path/filepath"
"strings"
"sync"
"time"
)

// Session manages the chat state and interacts with the LLM client.
type Session struct {
clientMu sync.RWMutex
client *client.Client
HistoryPath string
History []client.ChatMessage
Expand Down Expand Up @@ -161,7 +163,7 @@ func (s *Session) StartStream(ctx context.Context, extraBody map[string]any) (<-
req.Tools = s.GetToolDefinitions()
}

streamOut, streamErr := s.client.ChatCompletionStream(ctx, req)
streamOut, streamErr := s.Client().ChatCompletionStream(ctx, req)

go func() {
defer close(outCh)
Expand Down Expand Up @@ -229,7 +231,7 @@ func (s *Session) Impersonate(ctx context.Context) (string, error) {
N_Predict: 50,
}

resp, err := s.client.Completion(ctx, req)
resp, err := s.Client().Completion(ctx, req)
if err != nil {
return "", err
}
Expand Down Expand Up @@ -304,9 +306,19 @@ func (s *Session) saveAndNotify() error {
}

func (s *Session) Client() *client.Client {
s.clientMu.RLock()
defer s.clientMu.RUnlock()
return s.client
}

// SetClient replaces the client used for subsequent model requests.
// In-flight requests continue using the client they started with.
func (s *Session) SetClient(c *client.Client) {
s.clientMu.Lock()
defer s.clientMu.Unlock()
s.client = c
}

func (s *Session) IsLlamaCPP() bool {
return s.client.IsLlamaCPP()
return s.Client().IsLlamaCPP()
}
4 changes: 3 additions & 1 deletion internal/tui/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package tui

import (
"late/internal/common"
"late/internal/config"
"os"

"charm.land/bubbles/v2/filepicker"
Expand All @@ -13,7 +14,7 @@ import (
"charm.land/lipgloss/v2"
)

func NewModel(root common.Orchestrator, renderer *glamour.TermRenderer) Model {
func NewModel(root common.Orchestrator, renderer *glamour.TermRenderer, cfg *config.Config) Model {
ti := textarea.New()
ti.Placeholder = "Ask Late anything..."
ti.Focus()
Expand Down Expand Up @@ -84,6 +85,7 @@ func NewModel(root common.Orchestrator, renderer *glamour.TermRenderer) Model {
ShowCWD: true,
cachedRendererWidth: -1, // Force first creation
Pastes: make(map[string]string),
AppConfig: cfg,
}

fp := filepicker.New()
Expand Down
Loading