summaryrefslogtreecommitdiffhomepage
path: root/packages/tui/internal
diff options
context:
space:
mode:
Diffstat (limited to 'packages/tui/internal')
-rw-r--r--packages/tui/internal/app/app.go49
-rw-r--r--packages/tui/internal/components/chat/editor.go2
-rw-r--r--packages/tui/internal/components/chat/message.go52
-rw-r--r--packages/tui/internal/components/chat/messages.go29
-rw-r--r--packages/tui/internal/components/dialog/models.go8
-rw-r--r--packages/tui/internal/state/state.go4
6 files changed, 77 insertions, 67 deletions
diff --git a/packages/tui/internal/app/app.go b/packages/tui/internal/app/app.go
index 3917330e3..69900a0bd 100644
--- a/packages/tui/internal/app/app.go
+++ b/packages/tui/internal/app/app.go
@@ -23,7 +23,7 @@ type App struct {
Config *config.Config
Client *client.ClientWithResponses
Provider *client.ProviderInfo
- Model *client.ProviderModel
+ Model *client.ModelInfo
Session *client.SessionInfo
Messages []client.MessageInfo
Status status.Service
@@ -61,20 +61,25 @@ func New(ctx context.Context, version string, httpClient *client.ClientWithRespo
}
providers := []client.ProviderInfo{}
var defaultProvider *client.ProviderInfo
- var defaultModel *client.ProviderModel
-
- for i, provider := range providersResponse.JSON200.Providers {
- if i == 0 || provider.Id == "anthropic" {
- defaultProvider = &providersResponse.JSON200.Providers[i]
- if match, ok := providersResponse.JSON200.Default[provider.Id]; ok {
- model := defaultProvider.Models[match]
- defaultModel = &model
- } else {
- for _, model := range provider.Models {
- defaultModel = &model
- break
- }
- }
+ var defaultModel *client.ModelInfo
+
+ var anthropic *client.ProviderInfo
+ for _, provider := range providersResponse.JSON200.Providers {
+ if provider.Id == "anthropic" {
+ anthropic = &provider
+ }
+ }
+
+ // default to anthropic if available
+ if anthropic != nil {
+ defaultProvider = anthropic
+ defaultModel = getDefaultModel(providersResponse, *anthropic)
+ }
+
+ for _, provider := range providersResponse.JSON200.Providers {
+ if defaultProvider == nil || defaultModel == nil {
+ defaultProvider = &provider
+ defaultModel = getDefaultModel(providersResponse, provider)
}
providers = append(providers, provider)
}
@@ -91,7 +96,7 @@ func New(ctx context.Context, version string, httpClient *client.ClientWithRespo
}
var currentProvider *client.ProviderInfo
- var currentModel *client.ProviderModel
+ var currentModel *client.ModelInfo
for _, provider := range providers {
if provider.Id == appConfig.Provider {
currentProvider = &provider
@@ -121,6 +126,18 @@ func New(ctx context.Context, version string, httpClient *client.ClientWithRespo
return app, nil
}
+func getDefaultModel(response *client.PostProviderListResponse, provider client.ProviderInfo) *client.ModelInfo {
+ if match, ok := response.JSON200.Default[provider.Id]; ok {
+ model := provider.Models[match]
+ return &model
+ } else {
+ for _, model := range provider.Models {
+ return &model
+ }
+ }
+ return nil
+}
+
type Attachment struct {
FilePath string
FileName string
diff --git a/packages/tui/internal/components/chat/editor.go b/packages/tui/internal/components/chat/editor.go
index 8d977f5e0..f78bd1926 100644
--- a/packages/tui/internal/components/chat/editor.go
+++ b/packages/tui/internal/components/chat/editor.go
@@ -284,7 +284,7 @@ func (m *editorComponent) View() string {
model := ""
if m.app.Model != nil {
- model = base(*m.app.Model.Name) + muted(" • /model")
+ model = base(m.app.Model.Name) + muted(" • /model")
}
space := m.width - 2 - lipgloss.Width(model) - lipgloss.Width(hint)
diff --git a/packages/tui/internal/components/chat/message.go b/packages/tui/internal/components/chat/message.go
index 2c2cd03f0..82e749afe 100644
--- a/packages/tui/internal/components/chat/message.go
+++ b/packages/tui/internal/components/chat/message.go
@@ -212,7 +212,7 @@ func renderText(message client.MessageInfo, text string, author string) string {
func renderToolInvocation(
toolCall client.MessageToolInvocationToolCall,
result *string,
- metadata map[string]any,
+ metadata client.MessageInfo_Metadata_Tool_AdditionalProperties,
showResult bool,
) string {
ignoredTools := []string{"opencode_todoread"}
@@ -264,27 +264,26 @@ func renderToolInvocation(
body = *result
}
- if metadata["error"] != nil && metadata["message"] != nil {
- body = ""
- error = styles.BaseStyle().
- Foreground(t.Error()).
- Render(metadata["message"].(string))
- error = renderContentBlock(error, WithBorderColor(t.Error()), WithFullWidth(), WithPaddingTop(1), WithPaddingBottom(1))
+ if e, ok := metadata.Get("error"); ok && e.(bool) == true {
+ if m, ok := metadata.Get("message"); ok {
+ body = "" // don't show the body if there's an error
+ error = styles.BaseStyle().
+ Foreground(t.Error()).
+ Render(m.(string))
+ error = renderContentBlock(error, WithBorderColor(t.Error()), WithFullWidth(), WithPaddingTop(1), WithPaddingBottom(1))
+ }
}
elapsed := ""
- if metadata["time"] != nil {
- timeMap := metadata["time"].(map[string]any)
- start := timeMap["start"].(float64)
- end := timeMap["end"].(float64)
- durationMs := end - start
- duration := time.Duration(durationMs * float64(time.Millisecond))
- roundedDuration := time.Duration(duration.Round(time.Millisecond))
- if durationMs > 1000 {
- roundedDuration = time.Duration(duration.Round(time.Second))
- }
- elapsed = styles.Muted().Render(roundedDuration.String())
+ start := metadata.Time.Start
+ end := metadata.Time.End
+ durationMs := end - start
+ duration := time.Duration(durationMs * float32(time.Millisecond))
+ roundedDuration := time.Duration(duration.Round(time.Millisecond))
+ if durationMs > 1000 {
+ roundedDuration = time.Duration(duration.Round(time.Second))
}
+ elapsed = styles.Muted().Render(roundedDuration.String())
title := ""
switch toolCall.ToolName {
@@ -292,16 +291,16 @@ func renderToolInvocation(
toolArgs = renderArgs(&toolArgsMap, "filePath")
title = fmt.Sprintf("Read: %s %s", toolArgs, elapsed)
body = ""
- if metadata["preview"] != nil && toolArgsMap["filePath"] != nil {
+ if preview, ok := metadata.Get("preview"); ok && toolArgsMap["filePath"] != nil {
filename := toolArgsMap["filePath"].(string)
- body = metadata["preview"].(string)
+ body = preview.(string)
body = renderFile(filename, body, WithTruncate(6))
}
case "opencode_edit":
filename := toolArgsMap["filePath"].(string)
title = fmt.Sprintf("Edit: %s %s", relative(filename), elapsed)
- if metadata["diff"] != nil {
- patch := metadata["diff"].(string)
+ if d, ok := metadata.Get("diff"); ok {
+ patch := d.(string)
diffWidth := min(layout.Current.Viewport.Width, 120)
formattedDiff, _ := diff.FormatDiff(filename, patch, diff.WithTotalWidth(diffWidth))
body = strings.TrimSpace(formattedDiff)
@@ -322,9 +321,9 @@ func renderToolInvocation(
case "opencode_bash":
description := toolArgsMap["description"].(string)
title = fmt.Sprintf("Shell: %s %s", description, elapsed)
- if metadata["stdout"] != nil {
+ if stdout, ok := metadata.Get("stdout"); ok {
command := toolArgsMap["command"].(string)
- stdout := metadata["stdout"].(string)
+ stdout := stdout.(string)
body = fmt.Sprintf("```console\n> %s\n%s```", command, stdout)
body = toMarkdown(body, innerWidth, t.BackgroundSubtle())
body = renderContentBlock(body, WithFullWidth(), WithPaddingTop(1), WithPaddingBottom(1))
@@ -339,9 +338,10 @@ func renderToolInvocation(
body = renderContentBlock(body, WithFullWidth(), WithPaddingTop(1), WithPaddingBottom(1))
case "opencode_todowrite":
title = fmt.Sprintf("Planning... %s", elapsed)
- if finished && metadata["todos"] != nil {
+
+ if to, ok := metadata.Get("todos"); ok && finished {
body = ""
- todos := metadata["todos"].([]any)
+ todos := to.([]any)
for _, todo := range todos {
t := todo.(map[string]any)
content := t["content"].(string)
diff --git a/packages/tui/internal/components/chat/messages.go b/packages/tui/internal/components/chat/messages.go
index 3985ee0d7..ef35a3532 100644
--- a/packages/tui/internal/components/chat/messages.go
+++ b/packages/tui/internal/components/chat/messages.go
@@ -118,7 +118,6 @@ type blockType int
const (
none blockType = iota
- systemTextBlock
userTextBlock
assistantTextBlock
toolInvocationBlock
@@ -134,10 +133,6 @@ func (m *messagesComponent) renderView() {
blocks := make([]string, 0)
previousBlockType := none
for _, message := range m.app.Messages {
- if message.Role == client.System {
- continue // ignoring system messages for now
- }
-
var content string
var cached bool
@@ -174,15 +169,13 @@ func (m *messagesComponent) renderView() {
previousBlockType = userTextBlock
} else if message.Role == client.Assistant {
previousBlockType = assistantTextBlock
- } else if message.Role == client.System {
- previousBlockType = systemTextBlock
}
case client.MessagePartToolInvocation:
toolInvocationPart := part.(client.MessagePartToolInvocation)
toolCall, _ := toolInvocationPart.ToolInvocation.AsMessageToolInvocationToolCall()
- metadata := map[string]any{}
+ metadata := client.MessageInfo_Metadata_Tool_AdditionalProperties{}
if _, ok := message.Metadata.Tool[toolCall.ToolCallId]; ok {
- metadata = message.Metadata.Tool[toolCall.ToolCallId].(map[string]any)
+ metadata = message.Metadata.Tool[toolCall.ToolCallId]
}
var result *string
resultPart, resultError := toolInvocationPart.ToolInvocation.AsMessageToolInvocationToolResult()
@@ -215,14 +208,16 @@ func (m *messagesComponent) renderView() {
}
error := ""
- errorValue, _ := message.Metadata.Error.ValueByDiscriminator()
- switch errorValue.(type) {
- case client.UnknownError:
- clientError := errorValue.(client.UnknownError)
- error = clientError.Data.Message
- error = renderContentBlock(error, WithBorderColor(t.Error()), WithFullWidth(), WithPaddingTop(1), WithPaddingBottom(1))
- blocks = append(blocks, error)
- previousBlockType = errorBlock
+ if message.Metadata.Error != nil {
+ errorValue, _ := message.Metadata.Error.ValueByDiscriminator()
+ switch errorValue.(type) {
+ case client.UnknownError:
+ clientError := errorValue.(client.UnknownError)
+ error = clientError.Data.Message
+ error = renderContentBlock(error, WithBorderColor(t.Error()), WithFullWidth(), WithPaddingTop(1), WithPaddingBottom(1))
+ blocks = append(blocks, error)
+ previousBlockType = errorBlock
+ }
}
}
diff --git a/packages/tui/internal/components/dialog/models.go b/packages/tui/internal/components/dialog/models.go
index ca6561502..ed2ab3354 100644
--- a/packages/tui/internal/components/dialog/models.go
+++ b/packages/tui/internal/components/dialog/models.go
@@ -125,9 +125,9 @@ func (m *modelDialog) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
return m, nil
}
-func (m *modelDialog) models() []client.ProviderModel {
- models := slices.SortedFunc(maps.Values(m.provider.Models), func(a, b client.ProviderModel) int {
- return strings.Compare(*a.Name, *b.Name)
+func (m *modelDialog) models() []client.ModelInfo {
+ models := slices.SortedFunc(maps.Values(m.provider.Models), func(a, b client.ModelInfo) int {
+ return strings.Compare(a.Name, b.Name)
})
return models
}
@@ -205,7 +205,7 @@ func (m *modelDialog) View() string {
Foreground(t.BackgroundElement()).
Bold(true)
}
- modelItems = append(modelItems, itemStyle.Render(*models[i].Name))
+ modelItems = append(modelItems, itemStyle.Render(models[i].Name))
}
scrollIndicator := m.getScrollIndicators(maxDialogWidth)
diff --git a/packages/tui/internal/state/state.go b/packages/tui/internal/state/state.go
index d2cbf039c..c5322e7b2 100644
--- a/packages/tui/internal/state/state.go
+++ b/packages/tui/internal/state/state.go
@@ -7,7 +7,7 @@ import (
type SessionSelectedMsg = *client.SessionInfo
type ModelSelectedMsg struct {
Provider client.ProviderInfo
- Model client.ProviderModel
+ Model client.ModelInfo
}
type SessionClearedMsg struct{}
@@ -17,5 +17,3 @@ type CompactSessionMsg struct{}
type StateUpdatedMsg struct {
State map[string]any
}
-
-// TODO: store in CONFIG/tui.yaml