summaryrefslogtreecommitdiffhomepage
path: root/internal/llm/agent
diff options
context:
space:
mode:
authorKujtim Hoxha <[email protected]>2025-04-21 19:59:35 +0200
committerGitHub <[email protected]>2025-04-21 19:59:35 +0200
commitf33dff87725764af0b675b5e5b2e011b21c14c90 (patch)
tree4fe2c022305f13775f2cab3cdd80cd808259765b /internal/llm/agent
parent6b1c64bcc75b89c530294b6a2d4404682b435d56 (diff)
parent3a6a26981a8074b6ab0eaadb520db986e04799ff (diff)
downloadopencode-f33dff87725764af0b675b5e5b2e011b21c14c90.tar.gz
opencode-f33dff87725764af0b675b5e5b2e011b21c14c90.zip
Merge pull request #27 from kujtimiihoxha/opencode
OpenCode - Initial Implementation
Diffstat (limited to 'internal/llm/agent')
-rw-r--r--internal/llm/agent/agent-tool.go63
-rw-r--r--internal/llm/agent/agent.go816
-rw-r--r--internal/llm/agent/coder.go73
-rw-r--r--internal/llm/agent/mcp-tools.go19
-rw-r--r--internal/llm/agent/task.go46
-rw-r--r--internal/llm/agent/tools.go51
6 files changed, 494 insertions, 574 deletions
diff --git a/internal/llm/agent/agent-tool.go b/internal/llm/agent/agent-tool.go
index deb6aed60..be6e09a9b 100644
--- a/internal/llm/agent/agent-tool.go
+++ b/internal/llm/agent/agent-tool.go
@@ -5,14 +5,17 @@ import (
"encoding/json"
"fmt"
- "github.com/kujtimiihoxha/termai/internal/app"
- "github.com/kujtimiihoxha/termai/internal/llm/tools"
- "github.com/kujtimiihoxha/termai/internal/message"
+ "github.com/kujtimiihoxha/opencode/internal/config"
+ "github.com/kujtimiihoxha/opencode/internal/llm/tools"
+ "github.com/kujtimiihoxha/opencode/internal/lsp"
+ "github.com/kujtimiihoxha/opencode/internal/message"
+ "github.com/kujtimiihoxha/opencode/internal/session"
)
type agentTool struct {
- parentSessionID string
- app *app.App
+ sessions session.Service
+ messages message.Service
+ lspClients map[string]*lsp.Client
}
const (
@@ -46,57 +49,63 @@ func (b *agentTool) Run(ctx context.Context, call tools.ToolCall) (tools.ToolRes
return tools.NewTextErrorResponse("prompt is required"), nil
}
- agent, err := NewTaskAgent(b.app)
- if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error creating agent: %s", err)), nil
+ sessionID, messageID := tools.GetContextValues(ctx)
+ if sessionID == "" || messageID == "" {
+ return tools.ToolResponse{}, fmt.Errorf("session_id and message_id are required")
}
- session, err := b.app.Sessions.CreateTaskSession(call.ID, b.parentSessionID, "New Agent Session")
+ agent, err := NewAgent(config.AgentTask, b.sessions, b.messages, TaskAgentTools(b.lspClients))
if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error creating session: %s", err)), nil
+ return tools.ToolResponse{}, fmt.Errorf("error creating agent: %s", err)
}
- err = agent.Generate(ctx, session.ID, params.Prompt)
+ session, err := b.sessions.CreateTaskSession(ctx, call.ID, sessionID, "New Agent Session")
if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error generating agent: %s", err)), nil
+ return tools.ToolResponse{}, fmt.Errorf("error creating session: %s", err)
}
- messages, err := b.app.Messages.List(session.ID)
+ done, err := agent.Run(ctx, session.ID, params.Prompt)
if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error listing messages: %s", err)), nil
+ return tools.ToolResponse{}, fmt.Errorf("error generating agent: %s", err)
}
- if len(messages) == 0 {
- return tools.NewTextErrorResponse("no messages found"), nil
+ result := <-done
+ if result.Err() != nil {
+ return tools.ToolResponse{}, fmt.Errorf("error generating agent: %s", result.Err())
}
- response := messages[len(messages)-1]
+ response := result.Response()
if response.Role != message.Assistant {
- return tools.NewTextErrorResponse("no assistant message found"), nil
+ return tools.NewTextErrorResponse("no response"), nil
}
- updatedSession, err := b.app.Sessions.Get(session.ID)
+ updatedSession, err := b.sessions.Get(ctx, session.ID)
if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error: %s", err)), nil
+ return tools.ToolResponse{}, fmt.Errorf("error getting session: %s", err)
}
- parentSession, err := b.app.Sessions.Get(b.parentSessionID)
+ parentSession, err := b.sessions.Get(ctx, sessionID)
if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error: %s", err)), nil
+ return tools.ToolResponse{}, fmt.Errorf("error getting parent session: %s", err)
}
parentSession.Cost += updatedSession.Cost
parentSession.PromptTokens += updatedSession.PromptTokens
parentSession.CompletionTokens += updatedSession.CompletionTokens
- _, err = b.app.Sessions.Save(parentSession)
+ _, err = b.sessions.Save(ctx, parentSession)
if err != nil {
- return tools.NewTextErrorResponse(fmt.Sprintf("error: %s", err)), nil
+ return tools.ToolResponse{}, fmt.Errorf("error saving parent session: %s", err)
}
return tools.NewTextResponse(response.Content().String()), nil
}
-func NewAgentTool(parentSessionID string, app *app.App) tools.BaseTool {
+func NewAgentTool(
+ Sessions session.Service,
+ Messages message.Service,
+ LspClients map[string]*lsp.Client,
+) tools.BaseTool {
return &agentTool{
- parentSessionID: parentSessionID,
- app: app,
+ sessions: Sessions,
+ messages: Messages,
+ lspClients: LspClients,
}
}
diff --git a/internal/llm/agent/agent.go b/internal/llm/agent/agent.go
index 998dc1551..6c5808eab 100644
--- a/internal/llm/agent/agent.go
+++ b/internal/llm/agent/agent.go
@@ -7,30 +7,123 @@ import (
"strings"
"sync"
- "github.com/kujtimiihoxha/termai/internal/app"
- "github.com/kujtimiihoxha/termai/internal/config"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
- "github.com/kujtimiihoxha/termai/internal/llm/prompt"
- "github.com/kujtimiihoxha/termai/internal/llm/provider"
- "github.com/kujtimiihoxha/termai/internal/llm/tools"
- "github.com/kujtimiihoxha/termai/internal/logging"
- "github.com/kujtimiihoxha/termai/internal/message"
+ "github.com/kujtimiihoxha/opencode/internal/config"
+ "github.com/kujtimiihoxha/opencode/internal/llm/models"
+ "github.com/kujtimiihoxha/opencode/internal/llm/prompt"
+ "github.com/kujtimiihoxha/opencode/internal/llm/provider"
+ "github.com/kujtimiihoxha/opencode/internal/llm/tools"
+ "github.com/kujtimiihoxha/opencode/internal/logging"
+ "github.com/kujtimiihoxha/opencode/internal/message"
+ "github.com/kujtimiihoxha/opencode/internal/permission"
+ "github.com/kujtimiihoxha/opencode/internal/session"
)
-type Agent interface {
- Generate(ctx context.Context, sessionID string, content string) error
+// Common errors
+var (
+ ErrRequestCancelled = errors.New("request cancelled by user")
+ ErrSessionBusy = errors.New("session is currently processing another request")
+)
+
+type AgentEvent struct {
+ message message.Message
+ err error
+}
+
+func (e *AgentEvent) Err() error {
+ return e.err
+}
+
+func (e *AgentEvent) Response() message.Message {
+ return e.message
+}
+
+type Service interface {
+ Run(ctx context.Context, sessionID string, content string) (<-chan AgentEvent, error)
+ Cancel(sessionID string)
+ IsSessionBusy(sessionID string) bool
+ IsBusy() bool
}
type agent struct {
- *app.App
- model models.Model
- tools []tools.BaseTool
- agent provider.Provider
- titleGenerator provider.Provider
+ sessions session.Service
+ messages message.Service
+
+ tools []tools.BaseTool
+ provider provider.Provider
+
+ titleProvider provider.Provider
+
+ activeRequests sync.Map
+}
+
+func NewAgent(
+ agentName config.AgentName,
+ sessions session.Service,
+ messages message.Service,
+ agentTools []tools.BaseTool,
+) (Service, error) {
+ agentProvider, err := createAgentProvider(agentName)
+ if err != nil {
+ return nil, err
+ }
+ var titleProvider provider.Provider
+ // Only generate titles for the coder agent
+ if agentName == config.AgentCoder {
+ titleProvider, err = createAgentProvider(config.AgentTitle)
+ if err != nil {
+ return nil, err
+ }
+ }
+
+ agent := &agent{
+ provider: agentProvider,
+ messages: messages,
+ sessions: sessions,
+ tools: agentTools,
+ titleProvider: titleProvider,
+ activeRequests: sync.Map{},
+ }
+
+ return agent, nil
+}
+
+func (a *agent) Cancel(sessionID string) {
+ if cancelFunc, exists := a.activeRequests.LoadAndDelete(sessionID); exists {
+ if cancel, ok := cancelFunc.(context.CancelFunc); ok {
+ logging.InfoPersist(fmt.Sprintf("Request cancellation initiated for session: %s", sessionID))
+ cancel()
+ }
+ }
+}
+
+func (a *agent) IsBusy() bool {
+ busy := false
+ a.activeRequests.Range(func(key, value interface{}) bool {
+ if cancelFunc, ok := value.(context.CancelFunc); ok {
+ if cancelFunc != nil {
+ busy = true
+ return false // Stop iterating
+ }
+ }
+ return true // Continue iterating
+ })
+ return busy
+}
+
+func (a *agent) IsSessionBusy(sessionID string) bool {
+ _, busy := a.activeRequests.Load(sessionID)
+ return busy
}
-func (c *agent) handleTitleGeneration(ctx context.Context, sessionID, content string) {
- response, err := c.titleGenerator.SendMessages(
+func (a *agent) generateTitle(ctx context.Context, sessionID string, content string) error {
+ if a.titleProvider == nil {
+ return nil
+ }
+ session, err := a.sessions.Get(ctx, sessionID)
+ if err != nil {
+ return err
+ }
+ response, err := a.titleProvider.SendMessages(
ctx,
[]message.Message{
{
@@ -42,476 +135,357 @@ func (c *agent) handleTitleGeneration(ctx context.Context, sessionID, content st
},
},
},
- nil,
+ make([]tools.BaseTool, 0),
)
if err != nil {
- return
- }
-
- session, err := c.Sessions.Get(sessionID)
- if err != nil {
- return
- }
- if response.Content != "" {
- session.Title = response.Content
- session.Title = strings.TrimSpace(session.Title)
- session.Title = strings.ReplaceAll(session.Title, "\n", " ")
- c.Sessions.Save(session)
- }
-}
-
-func (c *agent) TrackUsage(sessionID string, model models.Model, usage provider.TokenUsage) error {
- session, err := c.Sessions.Get(sessionID)
- if err != nil {
return err
}
- cost := model.CostPer1MInCached/1e6*float64(usage.CacheCreationTokens) +
- model.CostPer1MOutCached/1e6*float64(usage.CacheReadTokens) +
- model.CostPer1MIn/1e6*float64(usage.InputTokens) +
- model.CostPer1MOut/1e6*float64(usage.OutputTokens)
-
- session.Cost += cost
- session.CompletionTokens += usage.OutputTokens
- session.PromptTokens += usage.InputTokens
+ title := strings.TrimSpace(strings.ReplaceAll(response.Content, "\n", " "))
+ if title == "" {
+ return nil
+ }
- _, err = c.Sessions.Save(session)
+ session.Title = title
+ _, err = a.sessions.Save(ctx, session)
return err
}
-func (c *agent) processEvent(
- sessionID string,
- assistantMsg *message.Message,
- event provider.ProviderEvent,
-) error {
- switch event.Type {
- case provider.EventThinkingDelta:
- assistantMsg.AppendReasoningContent(event.Content)
- return c.Messages.Update(*assistantMsg)
- case provider.EventContentDelta:
- assistantMsg.AppendContent(event.Content)
- return c.Messages.Update(*assistantMsg)
- case provider.EventError:
- if errors.Is(event.Error, context.Canceled) {
- return nil
- }
- logging.ErrorPersist(event.Error.Error())
- return event.Error
- case provider.EventWarning:
- logging.WarnPersist(event.Info)
- return nil
- case provider.EventInfo:
- logging.InfoPersist(event.Info)
- case provider.EventComplete:
- assistantMsg.SetToolCalls(event.Response.ToolCalls)
- assistantMsg.AddFinish(event.Response.FinishReason)
- err := c.Messages.Update(*assistantMsg)
- if err != nil {
- return err
- }
- return c.TrackUsage(sessionID, c.model, event.Response.Usage)
+func (a *agent) err(err error) AgentEvent {
+ return AgentEvent{
+ err: err,
}
-
- return nil
}
-func (c *agent) ExecuteTools(ctx context.Context, toolCalls []message.ToolCall, tls []tools.BaseTool) ([]message.ToolResult, error) {
- var wg sync.WaitGroup
- toolResults := make([]message.ToolResult, len(toolCalls))
- mutex := &sync.Mutex{}
- errChan := make(chan error, 1)
-
- // Create a child context that can be canceled
- ctx, cancel := context.WithCancel(ctx)
- defer cancel()
-
- for i, tc := range toolCalls {
- wg.Add(1)
- go func(index int, toolCall message.ToolCall) {
- defer wg.Done()
-
- // Check if context is already canceled
- select {
- case <-ctx.Done():
- mutex.Lock()
- toolResults[index] = message.ToolResult{
- ToolCallID: toolCall.ID,
- Content: "Tool execution canceled",
- IsError: true,
- }
- mutex.Unlock()
-
- // Send cancellation error to error channel if it's empty
- select {
- case errChan <- ctx.Err():
- default:
- }
- return
- default:
- }
-
- response := ""
- isError := false
- found := false
-
- for _, tool := range tls {
- if tool.Info().Name == toolCall.Name {
- found = true
- toolResult, toolErr := tool.Run(ctx, tools.ToolCall{
- ID: toolCall.ID,
- Name: toolCall.Name,
- Input: toolCall.Input,
- })
-
- if toolErr != nil {
- if errors.Is(toolErr, context.Canceled) {
- response = "Tool execution canceled"
-
- // Send cancellation error to error channel if it's empty
- select {
- case errChan <- ctx.Err():
- default:
- }
- } else {
- response = fmt.Sprintf("error running tool: %s", toolErr)
- }
- isError = true
- } else {
- response = toolResult.Content
- isError = toolResult.IsError
- }
- break
- }
- }
-
- if !found {
- response = fmt.Sprintf("tool not found: %s", toolCall.Name)
- isError = true
- }
-
- mutex.Lock()
- defer mutex.Unlock()
-
- toolResults[index] = message.ToolResult{
- ToolCallID: toolCall.ID,
- Content: response,
- IsError: isError,
- }
- }(i, tc)
+func (a *agent) Run(ctx context.Context, sessionID string, content string) (<-chan AgentEvent, error) {
+ events := make(chan AgentEvent)
+ if a.IsSessionBusy(sessionID) {
+ return nil, ErrSessionBusy
}
- // Wait for all goroutines to finish or context to be canceled
- done := make(chan struct{})
- go func() {
- wg.Wait()
- close(done)
- }()
+ genCtx, cancel := context.WithCancel(ctx)
- select {
- case <-done:
- // All tools completed successfully
- case err := <-errChan:
- // One of the tools encountered a cancellation
- return toolResults, err
- case <-ctx.Done():
- // Context was canceled externally
- return toolResults, ctx.Err()
- }
+ a.activeRequests.Store(sessionID, cancel)
+ go func() {
+ logging.Debug("Request started", "sessionID", sessionID)
+ defer logging.RecoverPanic("agent.Run", func() {
+ events <- a.err(fmt.Errorf("panic while running the agent"))
+ })
- return toolResults, nil
+ result := a.processGeneration(genCtx, sessionID, content)
+ if result.Err() != nil && !errors.Is(result.Err(), ErrRequestCancelled) && !errors.Is(result.Err(), context.Canceled) {
+ logging.ErrorPersist(fmt.Sprintf("Generation error for session %s: %v", sessionID, result))
+ }
+ logging.Debug("Request completed", "sessionID", sessionID)
+ a.activeRequests.Delete(sessionID)
+ cancel()
+ events <- result
+ close(events)
+ }()
+ return events, nil
}
-func (c *agent) handleToolExecution(
- ctx context.Context,
- assistantMsg message.Message,
-) (*message.Message, error) {
- if len(assistantMsg.ToolCalls()) == 0 {
- return nil, nil
- }
-
- toolResults, err := c.ExecuteTools(ctx, assistantMsg.ToolCalls(), c.tools)
+func (a *agent) processGeneration(ctx context.Context, sessionID, content string) AgentEvent {
+ // List existing messages; if none, start title generation asynchronously.
+ msgs, err := a.messages.List(ctx, sessionID)
if err != nil {
- return nil, err
+ return a.err(fmt.Errorf("failed to list messages: %w", err))
}
- parts := make([]message.ContentPart, 0)
- for _, toolResult := range toolResults {
- parts = append(parts, toolResult)
+ if len(msgs) == 0 {
+ go func() {
+ defer logging.RecoverPanic("agent.Run", func() {
+ logging.ErrorPersist("panic while generating title")
+ })
+ titleErr := a.generateTitle(context.Background(), sessionID, content)
+ if titleErr != nil {
+ logging.ErrorPersist(fmt.Sprintf("failed to generate title: %v", titleErr))
+ }
+ }()
}
- msg, err := c.Messages.Create(assistantMsg.SessionID, message.CreateMessageParams{
- Role: message.Tool,
- Parts: parts,
- })
-
- return &msg, err
-}
-func (c *agent) generate(ctx context.Context, sessionID string, content string) error {
- messages, err := c.Messages.List(sessionID)
+ userMsg, err := a.createUserMessage(ctx, sessionID, content)
if err != nil {
- return err
+ return a.err(fmt.Errorf("failed to create user message: %w", err))
}
- if len(messages) == 0 {
- go c.handleTitleGeneration(ctx, sessionID, content)
+ // Append the new user message to the conversation history.
+ msgHistory := append(msgs, userMsg)
+ for {
+ // Check for cancellation before each iteration
+ select {
+ case <-ctx.Done():
+ return a.err(ctx.Err())
+ default:
+ // Continue processing
+ }
+ agentMessage, toolResults, err := a.streamAndHandleEvents(ctx, sessionID, msgHistory)
+ if err != nil {
+ if errors.Is(err, context.Canceled) {
+ agentMessage.AddFinish(message.FinishReasonCanceled)
+ a.messages.Update(context.Background(), agentMessage)
+ return a.err(ErrRequestCancelled)
+ }
+ return a.err(fmt.Errorf("failed to process events: %w", err))
+ }
+ logging.Info("Result", "message", agentMessage.FinishReason(), "toolResults", toolResults)
+ if (agentMessage.FinishReason() == message.FinishReasonToolUse) && toolResults != nil {
+ // We are not done, we need to respond with the tool response
+ msgHistory = append(msgHistory, agentMessage, *toolResults)
+ continue
+ }
+ return AgentEvent{
+ message: agentMessage,
+ }
}
+}
- userMsg, err := c.Messages.Create(sessionID, message.CreateMessageParams{
+func (a *agent) createUserMessage(ctx context.Context, sessionID, content string) (message.Message, error) {
+ return a.messages.Create(ctx, sessionID, message.CreateMessageParams{
Role: message.User,
Parts: []message.ContentPart{
- message.TextContent{
- Text: content,
- },
+ message.TextContent{Text: content},
},
})
+}
+
+func (a *agent) streamAndHandleEvents(ctx context.Context, sessionID string, msgHistory []message.Message) (message.Message, *message.Message, error) {
+ eventChan := a.provider.StreamResponse(ctx, msgHistory, a.tools)
+
+ assistantMsg, err := a.messages.Create(ctx, sessionID, message.CreateMessageParams{
+ Role: message.Assistant,
+ Parts: []message.ContentPart{},
+ Model: a.provider.Model().ID,
+ })
if err != nil {
- return err
+ return assistantMsg, nil, fmt.Errorf("failed to create assistant message: %w", err)
}
- messages = append(messages, userMsg)
- for {
+ // Add the session and message ID into the context if needed by tools.
+ ctx = context.WithValue(ctx, tools.MessageIDContextKey, assistantMsg.ID)
+ ctx = context.WithValue(ctx, tools.SessionIDContextKey, sessionID)
+
+ // Process each event in the stream.
+ for event := range eventChan {
+ if processErr := a.processEvent(ctx, sessionID, &assistantMsg, event); processErr != nil {
+ a.finishMessage(ctx, &assistantMsg, message.FinishReasonCanceled)
+ return assistantMsg, nil, processErr
+ }
+ if ctx.Err() != nil {
+ a.finishMessage(context.Background(), &assistantMsg, message.FinishReasonCanceled)
+ return assistantMsg, nil, ctx.Err()
+ }
+ }
+
+ toolResults := make([]message.ToolResult, len(assistantMsg.ToolCalls()))
+ toolCalls := assistantMsg.ToolCalls()
+ for i, toolCall := range toolCalls {
select {
case <-ctx.Done():
- assistantMsg, err := c.Messages.Create(sessionID, message.CreateMessageParams{
- Role: message.Assistant,
- Parts: []message.ContentPart{},
- })
- if err != nil {
- return err
+ a.finishMessage(context.Background(), &assistantMsg, message.FinishReasonCanceled)
+ // Make all future tool calls cancelled
+ for j := i; j < len(toolCalls); j++ {
+ toolResults[j] = message.ToolResult{
+ ToolCallID: toolCalls[j].ID,
+ Content: "Tool execution canceled by user",
+ IsError: true,
+ }
}
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
- return context.Canceled
+ goto out
default:
// Continue processing
- }
-
- eventChan, err := c.agent.StreamResponse(ctx, messages, c.tools)
- if err != nil {
- if errors.Is(err, context.Canceled) {
- assistantMsg, err := c.Messages.Create(sessionID, message.CreateMessageParams{
- Role: message.Assistant,
- Parts: []message.ContentPart{},
- })
- if err != nil {
- return err
+ var tool tools.BaseTool
+ for _, availableTools := range a.tools {
+ if availableTools.Info().Name == toolCall.Name {
+ tool = availableTools
}
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
- return context.Canceled
}
- return err
- }
- assistantMsg, err := c.Messages.Create(sessionID, message.CreateMessageParams{
- Role: message.Assistant,
- Parts: []message.ContentPart{},
- })
- if err != nil {
- return err
- }
- for event := range eventChan {
- err = c.processEvent(sessionID, &assistantMsg, event)
- if err != nil {
- if errors.Is(err, context.Canceled) {
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
- return context.Canceled
+ // Tool not found
+ if tool == nil {
+ toolResults[i] = message.ToolResult{
+ ToolCallID: toolCall.ID,
+ Content: fmt.Sprintf("Tool not found: %s", toolCall.Name),
+ IsError: true,
}
- assistantMsg.AddFinish("error:" + err.Error())
- c.Messages.Update(assistantMsg)
- return err
+ continue
}
- select {
- case <-ctx.Done():
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
- return context.Canceled
- default:
+ toolResult, toolErr := tool.Run(ctx, tools.ToolCall{
+ ID: toolCall.ID,
+ Name: toolCall.Name,
+ Input: toolCall.Input,
+ })
+ if toolErr != nil {
+ if errors.Is(toolErr, permission.ErrorPermissionDenied) {
+ toolResults[i] = message.ToolResult{
+ ToolCallID: toolCall.ID,
+ Content: "Permission denied",
+ IsError: true,
+ }
+ for j := i + 1; j < len(toolCalls); j++ {
+ toolResults[j] = message.ToolResult{
+ ToolCallID: toolCalls[j].ID,
+ Content: "Tool execution canceled by user",
+ IsError: true,
+ }
+ }
+ a.finishMessage(ctx, &assistantMsg, message.FinishReasonPermissionDenied)
+ break
+ }
}
- }
-
- // Check for context cancellation before tool execution
- select {
- case <-ctx.Done():
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
- return context.Canceled
- default:
- // Continue processing
- }
-
- msg, err := c.handleToolExecution(ctx, assistantMsg)
- if err != nil {
- if errors.Is(err, context.Canceled) {
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
- return context.Canceled
+ toolResults[i] = message.ToolResult{
+ ToolCallID: toolCall.ID,
+ Content: toolResult.Content,
+ Metadata: toolResult.Metadata,
+ IsError: toolResult.IsError,
}
- return err
}
+ }
+out:
+ if len(toolResults) == 0 {
+ return assistantMsg, nil, nil
+ }
+ parts := make([]message.ContentPart, 0)
+ for _, tr := range toolResults {
+ parts = append(parts, tr)
+ }
+ msg, err := a.messages.Create(context.Background(), assistantMsg.SessionID, message.CreateMessageParams{
+ Role: message.Tool,
+ Parts: parts,
+ })
+ if err != nil {
+ return assistantMsg, nil, fmt.Errorf("failed to create cancelled tool message: %w", err)
+ }
- c.Messages.Update(assistantMsg)
+ return assistantMsg, &msg, err
+}
- if len(assistantMsg.ToolCalls()) == 0 {
- break
- }
+func (a *agent) finishMessage(ctx context.Context, msg *message.Message, finishReson message.FinishReason) {
+ msg.AddFinish(finishReson)
+ _ = a.messages.Update(ctx, *msg)
+}
- messages = append(messages, assistantMsg)
- if msg != nil {
- messages = append(messages, *msg)
- }
+func (a *agent) processEvent(ctx context.Context, sessionID string, assistantMsg *message.Message, event provider.ProviderEvent) error {
+ select {
+ case <-ctx.Done():
+ return ctx.Err()
+ default:
+ // Continue processing.
+ }
- // Check for context cancellation after tool execution
- select {
- case <-ctx.Done():
- assistantMsg.AddFinish("canceled")
- c.Messages.Update(assistantMsg)
+ switch event.Type {
+ case provider.EventThinkingDelta:
+ assistantMsg.AppendReasoningContent(event.Content)
+ return a.messages.Update(ctx, *assistantMsg)
+ case provider.EventContentDelta:
+ assistantMsg.AppendContent(event.Content)
+ return a.messages.Update(ctx, *assistantMsg)
+ case provider.EventToolUseStart:
+ assistantMsg.AddToolCall(*event.ToolCall)
+ return a.messages.Update(ctx, *assistantMsg)
+ // TODO: see how to handle this
+ // case provider.EventToolUseDelta:
+ // tm := time.Unix(assistantMsg.UpdatedAt, 0)
+ // assistantMsg.AppendToolCallInput(event.ToolCall.ID, event.ToolCall.Input)
+ // if time.Since(tm) > 1000*time.Millisecond {
+ // err := a.messages.Update(ctx, *assistantMsg)
+ // assistantMsg.UpdatedAt = time.Now().Unix()
+ // return err
+ // }
+ case provider.EventToolUseStop:
+ assistantMsg.FinishToolCall(event.ToolCall.ID)
+ return a.messages.Update(ctx, *assistantMsg)
+ case provider.EventError:
+ if errors.Is(event.Error, context.Canceled) {
+ logging.InfoPersist(fmt.Sprintf("Event processing canceled for session: %s", sessionID))
return context.Canceled
- default:
- // Continue processing
}
+ logging.ErrorPersist(event.Error.Error())
+ return event.Error
+ case provider.EventComplete:
+ assistantMsg.SetToolCalls(event.Response.ToolCalls)
+ assistantMsg.AddFinish(event.Response.FinishReason)
+ if err := a.messages.Update(ctx, *assistantMsg); err != nil {
+ return fmt.Errorf("failed to update message: %w", err)
+ }
+ return a.TrackUsage(ctx, sessionID, a.provider.Model(), event.Response.Usage)
}
+
return nil
}
-func getAgentProviders(ctx context.Context, model models.Model) (provider.Provider, provider.Provider, error) {
- maxTokens := config.Get().Model.CoderMaxTokens
-
- providerConfig, ok := config.Get().Providers[model.Provider]
- if !ok || !providerConfig.Enabled {
- return nil, nil, errors.New("provider is not enabled")
+func (a *agent) TrackUsage(ctx context.Context, sessionID string, model models.Model, usage provider.TokenUsage) error {
+ sess, err := a.sessions.Get(ctx, sessionID)
+ if err != nil {
+ return fmt.Errorf("failed to get session: %w", err)
}
- var agentProvider provider.Provider
- var titleGenerator provider.Provider
- switch model.Provider {
- case models.ProviderOpenAI:
- var err error
- agentProvider, err = provider.NewOpenAIProvider(
- provider.WithOpenAISystemMessage(
- prompt.CoderOpenAISystemPrompt(),
- ),
- provider.WithOpenAIMaxTokens(maxTokens),
- provider.WithOpenAIModel(model),
- provider.WithOpenAIKey(providerConfig.APIKey),
- )
- if err != nil {
- return nil, nil, err
- }
- titleGenerator, err = provider.NewOpenAIProvider(
- provider.WithOpenAISystemMessage(
- prompt.TitlePrompt(),
- ),
- provider.WithOpenAIMaxTokens(80),
- provider.WithOpenAIModel(model),
- provider.WithOpenAIKey(providerConfig.APIKey),
- )
- if err != nil {
- return nil, nil, err
- }
- case models.ProviderAnthropic:
- var err error
- agentProvider, err = provider.NewAnthropicProvider(
- provider.WithAnthropicSystemMessage(
- prompt.CoderAnthropicSystemPrompt(),
- ),
- provider.WithAnthropicMaxTokens(maxTokens),
- provider.WithAnthropicKey(providerConfig.APIKey),
- provider.WithAnthropicModel(model),
- )
- if err != nil {
- return nil, nil, err
- }
- titleGenerator, err = provider.NewAnthropicProvider(
- provider.WithAnthropicSystemMessage(
- prompt.TitlePrompt(),
- ),
- provider.WithAnthropicMaxTokens(80),
- provider.WithAnthropicKey(providerConfig.APIKey),
- provider.WithAnthropicModel(model),
- )
- if err != nil {
- return nil, nil, err
- }
+ cost := model.CostPer1MInCached/1e6*float64(usage.CacheCreationTokens) +
+ model.CostPer1MOutCached/1e6*float64(usage.CacheReadTokens) +
+ model.CostPer1MIn/1e6*float64(usage.InputTokens) +
+ model.CostPer1MOut/1e6*float64(usage.OutputTokens)
- case models.ProviderGemini:
- var err error
- agentProvider, err = provider.NewGeminiProvider(
- ctx,
- provider.WithGeminiSystemMessage(
- prompt.CoderOpenAISystemPrompt(),
- ),
- provider.WithGeminiMaxTokens(int32(maxTokens)),
- provider.WithGeminiKey(providerConfig.APIKey),
- provider.WithGeminiModel(model),
- )
- if err != nil {
- return nil, nil, err
- }
- titleGenerator, err = provider.NewGeminiProvider(
- ctx,
- provider.WithGeminiSystemMessage(
- prompt.TitlePrompt(),
- ),
- provider.WithGeminiMaxTokens(80),
- provider.WithGeminiKey(providerConfig.APIKey),
- provider.WithGeminiModel(model),
- )
- if err != nil {
- return nil, nil, err
- }
- case models.ProviderGROQ:
- var err error
- agentProvider, err = provider.NewOpenAIProvider(
- provider.WithOpenAISystemMessage(
- prompt.CoderAnthropicSystemPrompt(),
- ),
- provider.WithOpenAIMaxTokens(maxTokens),
- provider.WithOpenAIModel(model),
- provider.WithOpenAIKey(providerConfig.APIKey),
- provider.WithOpenAIBaseURL("https://api.groq.com/openai/v1"),
- )
- if err != nil {
- return nil, nil, err
- }
- titleGenerator, err = provider.NewOpenAIProvider(
- provider.WithOpenAISystemMessage(
- prompt.TitlePrompt(),
- ),
- provider.WithOpenAIMaxTokens(80),
- provider.WithOpenAIModel(model),
- provider.WithOpenAIKey(providerConfig.APIKey),
- provider.WithOpenAIBaseURL("https://api.groq.com/openai/v1"),
- )
- if err != nil {
- return nil, nil, err
- }
+ sess.Cost += cost
+ sess.CompletionTokens += usage.OutputTokens
+ sess.PromptTokens += usage.InputTokens
+
+ _, err = a.sessions.Save(ctx, sess)
+ if err != nil {
+ return fmt.Errorf("failed to save session: %w", err)
+ }
+ return nil
+}
+
+func createAgentProvider(agentName config.AgentName) (provider.Provider, error) {
+ cfg := config.Get()
+ agentConfig, ok := cfg.Agents[agentName]
+ if !ok {
+ return nil, fmt.Errorf("agent %s not found", agentName)
+ }
+ model, ok := models.SupportedModels[agentConfig.Model]
+ if !ok {
+ return nil, fmt.Errorf("model %s not supported", agentConfig.Model)
+ }
- case models.ProviderBedrock:
- var err error
- agentProvider, err = provider.NewBedrockProvider(
- provider.WithBedrockSystemMessage(
- prompt.CoderAnthropicSystemPrompt(),
+ providerCfg, ok := cfg.Providers[model.Provider]
+ if !ok {
+ return nil, fmt.Errorf("provider %s not supported", model.Provider)
+ }
+ if providerCfg.Disabled {
+ return nil, fmt.Errorf("provider %s is not enabled", model.Provider)
+ }
+ maxTokens := model.DefaultMaxTokens
+ if agentConfig.MaxTokens > 0 {
+ maxTokens = agentConfig.MaxTokens
+ }
+ opts := []provider.ProviderClientOption{
+ provider.WithAPIKey(providerCfg.APIKey),
+ provider.WithModel(model),
+ provider.WithSystemMessage(prompt.GetAgentPrompt(agentName, model.Provider)),
+ provider.WithMaxTokens(maxTokens),
+ }
+ if model.Provider == models.ProviderOpenAI && model.CanReason {
+ opts = append(
+ opts,
+ provider.WithOpenAIOptions(
+ provider.WithReasoningEffort(agentConfig.ReasoningEffort),
),
- provider.WithBedrockMaxTokens(maxTokens),
- provider.WithBedrockModel(model),
)
- if err != nil {
- return nil, nil, err
- }
- titleGenerator, err = provider.NewBedrockProvider(
- provider.WithBedrockSystemMessage(
- prompt.TitlePrompt(),
+ } else if model.Provider == models.ProviderAnthropic && model.CanReason && agentName == config.AgentCoder {
+ opts = append(
+ opts,
+ provider.WithAnthropicOptions(
+ provider.WithAnthropicShouldThinkFn(provider.DefaultShouldThinkFn),
),
- provider.WithBedrockMaxTokens(maxTokens),
- provider.WithBedrockModel(model),
)
- if err != nil {
- return nil, nil, err
- }
-
+ }
+ agentProvider, err := provider.NewProvider(
+ model.Provider,
+ opts...,
+ )
+ if err != nil {
+ return nil, fmt.Errorf("could not create provider: %v", err)
}
- return agentProvider, titleGenerator, nil
+ return agentProvider, nil
}
diff --git a/internal/llm/agent/coder.go b/internal/llm/agent/coder.go
deleted file mode 100644
index 5deff05a8..000000000
--- a/internal/llm/agent/coder.go
+++ /dev/null
@@ -1,73 +0,0 @@
-package agent
-
-import (
- "context"
- "errors"
-
- "github.com/kujtimiihoxha/termai/internal/app"
- "github.com/kujtimiihoxha/termai/internal/config"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
- "github.com/kujtimiihoxha/termai/internal/llm/tools"
-)
-
-type coderAgent struct {
- *agent
-}
-
-func (c *coderAgent) setAgentTool(sessionID string) {
- inx := -1
- for i, tool := range c.tools {
- if tool.Info().Name == AgentToolName {
- inx = i
- break
- }
- }
- if inx == -1 {
- c.tools = append(c.tools, NewAgentTool(sessionID, c.App))
- } else {
- c.tools[inx] = NewAgentTool(sessionID, c.App)
- }
-}
-
-func (c *coderAgent) Generate(ctx context.Context, sessionID string, content string) error {
- c.setAgentTool(sessionID)
- return c.generate(ctx, sessionID, content)
-}
-
-func NewCoderAgent(app *app.App) (Agent, error) {
- model, ok := models.SupportedModels[config.Get().Model.Coder]
- if !ok {
- return nil, errors.New("model not supported")
- }
-
- agentProvider, titleGenerator, err := getAgentProviders(app.Context, model)
- if err != nil {
- return nil, err
- }
-
- otherTools := GetMcpTools(app.Context, app.Permissions)
- if len(app.LSPClients) > 0 {
- otherTools = append(otherTools, tools.NewDiagnosticsTool(app.LSPClients))
- }
- return &coderAgent{
- agent: &agent{
- App: app,
- tools: append(
- []tools.BaseTool{
- tools.NewBashTool(app.Permissions),
- tools.NewEditTool(app.LSPClients, app.Permissions),
- tools.NewFetchTool(app.Permissions),
- tools.NewGlobTool(),
- tools.NewGrepTool(),
- tools.NewLsTool(),
- tools.NewSourcegraphTool(),
- tools.NewViewTool(app.LSPClients),
- tools.NewWriteTool(app.LSPClients, app.Permissions),
- }, otherTools...,
- ),
- model: model,
- agent: agentProvider,
- titleGenerator: titleGenerator,
- },
- }, nil
-}
diff --git a/internal/llm/agent/mcp-tools.go b/internal/llm/agent/mcp-tools.go
index b1c97b512..53aada33f 100644
--- a/internal/llm/agent/mcp-tools.go
+++ b/internal/llm/agent/mcp-tools.go
@@ -5,11 +5,11 @@ import (
"encoding/json"
"fmt"
- "github.com/kujtimiihoxha/termai/internal/config"
- "github.com/kujtimiihoxha/termai/internal/llm/tools"
- "github.com/kujtimiihoxha/termai/internal/logging"
- "github.com/kujtimiihoxha/termai/internal/permission"
- "github.com/kujtimiihoxha/termai/internal/version"
+ "github.com/kujtimiihoxha/opencode/internal/config"
+ "github.com/kujtimiihoxha/opencode/internal/llm/tools"
+ "github.com/kujtimiihoxha/opencode/internal/logging"
+ "github.com/kujtimiihoxha/opencode/internal/permission"
+ "github.com/kujtimiihoxha/opencode/internal/version"
"github.com/mark3labs/mcp-go/client"
"github.com/mark3labs/mcp-go/mcp"
@@ -46,7 +46,7 @@ func runTool(ctx context.Context, c MCPClient, toolName string, input string) (t
initRequest := mcp.InitializeRequest{}
initRequest.Params.ProtocolVersion = mcp.LATEST_PROTOCOL_VERSION
initRequest.Params.ClientInfo = mcp.Implementation{
- Name: "termai",
+ Name: "OpenCode",
Version: version.Version,
}
@@ -80,9 +80,14 @@ func runTool(ctx context.Context, c MCPClient, toolName string, input string) (t
}
func (b *mcpTool) Run(ctx context.Context, params tools.ToolCall) (tools.ToolResponse, error) {
+ sessionID, messageID := tools.GetContextValues(ctx)
+ if sessionID == "" || messageID == "" {
+ return tools.ToolResponse{}, fmt.Errorf("session ID and message ID are required for creating a new file")
+ }
permissionDescription := fmt.Sprintf("execute %s with the following parameters: %s", b.Info().Name, params.Input)
p := b.permissions.Request(
permission.CreatePermissionRequest{
+ SessionID: sessionID,
Path: config.WorkingDirectory(),
ToolName: b.Info().Name,
Action: "execute",
@@ -135,7 +140,7 @@ func getTools(ctx context.Context, name string, m config.MCPServer, permissions
initRequest := mcp.InitializeRequest{}
initRequest.Params.ProtocolVersion = mcp.LATEST_PROTOCOL_VERSION
initRequest.Params.ClientInfo = mcp.Implementation{
- Name: "termai",
+ Name: "OpenCode",
Version: version.Version,
}
diff --git a/internal/llm/agent/task.go b/internal/llm/agent/task.go
deleted file mode 100644
index 034e93460..000000000
--- a/internal/llm/agent/task.go
+++ /dev/null
@@ -1,46 +0,0 @@
-package agent
-
-import (
- "context"
- "errors"
-
- "github.com/kujtimiihoxha/termai/internal/app"
- "github.com/kujtimiihoxha/termai/internal/config"
- "github.com/kujtimiihoxha/termai/internal/llm/models"
- "github.com/kujtimiihoxha/termai/internal/llm/tools"
-)
-
-type taskAgent struct {
- *agent
-}
-
-func (c *taskAgent) Generate(ctx context.Context, sessionID string, content string) error {
- return c.generate(ctx, sessionID, content)
-}
-
-func NewTaskAgent(app *app.App) (Agent, error) {
- model, ok := models.SupportedModels[config.Get().Model.Coder]
- if !ok {
- return nil, errors.New("model not supported")
- }
-
- agentProvider, titleGenerator, err := getAgentProviders(app.Context, model)
- if err != nil {
- return nil, err
- }
- return &taskAgent{
- agent: &agent{
- App: app,
- tools: []tools.BaseTool{
- tools.NewGlobTool(),
- tools.NewGrepTool(),
- tools.NewLsTool(),
- tools.NewSourcegraphTool(),
- tools.NewViewTool(app.LSPClients),
- },
- model: model,
- agent: agentProvider,
- titleGenerator: titleGenerator,
- },
- }, nil
-}
diff --git a/internal/llm/agent/tools.go b/internal/llm/agent/tools.go
new file mode 100644
index 000000000..b2e6816d5
--- /dev/null
+++ b/internal/llm/agent/tools.go
@@ -0,0 +1,51 @@
+package agent
+
+import (
+ "context"
+
+ "github.com/kujtimiihoxha/opencode/internal/history"
+ "github.com/kujtimiihoxha/opencode/internal/llm/tools"
+ "github.com/kujtimiihoxha/opencode/internal/lsp"
+ "github.com/kujtimiihoxha/opencode/internal/message"
+ "github.com/kujtimiihoxha/opencode/internal/permission"
+ "github.com/kujtimiihoxha/opencode/internal/session"
+)
+
+func CoderAgentTools(
+ permissions permission.Service,
+ sessions session.Service,
+ messages message.Service,
+ history history.Service,
+ lspClients map[string]*lsp.Client,
+) []tools.BaseTool {
+ ctx := context.Background()
+ otherTools := GetMcpTools(ctx, permissions)
+ if len(lspClients) > 0 {
+ otherTools = append(otherTools, tools.NewDiagnosticsTool(lspClients))
+ }
+ return append(
+ []tools.BaseTool{
+ tools.NewBashTool(permissions),
+ tools.NewEditTool(lspClients, permissions, history),
+ tools.NewFetchTool(permissions),
+ tools.NewGlobTool(),
+ tools.NewGrepTool(),
+ tools.NewLsTool(),
+ tools.NewSourcegraphTool(),
+ tools.NewViewTool(lspClients),
+ tools.NewPatchTool(lspClients, permissions, history),
+ tools.NewWriteTool(lspClients, permissions, history),
+ NewAgentTool(sessions, messages, lspClients),
+ }, otherTools...,
+ )
+}
+
+func TaskAgentTools(lspClients map[string]*lsp.Client) []tools.BaseTool {
+ return []tools.BaseTool{
+ tools.NewGlobTool(),
+ tools.NewGrepTool(),
+ tools.NewLsTool(),
+ tools.NewSourcegraphTool(),
+ tools.NewViewTool(lspClients),
+ }
+}