summaryrefslogtreecommitdiffhomepage
path: root/internal/message
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/message
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/message')
-rw-r--r--internal/message/content.go88
-rw-r--r--internal/message/message.go67
2 files changed, 123 insertions, 32 deletions
diff --git a/internal/message/content.go b/internal/message/content.go
index 2604cd68a..beebe354e 100644
--- a/internal/message/content.go
+++ b/internal/message/content.go
@@ -2,6 +2,10 @@ package message
import (
"encoding/base64"
+ "slices"
+ "time"
+
+ "github.com/kujtimiihoxha/opencode/internal/llm/models"
)
type MessageRole string
@@ -13,6 +17,20 @@ const (
Tool MessageRole = "tool"
)
+type FinishReason string
+
+const (
+ FinishReasonEndTurn FinishReason = "end_turn"
+ FinishReasonMaxTokens FinishReason = "max_tokens"
+ FinishReasonToolUse FinishReason = "tool_use"
+ FinishReasonCanceled FinishReason = "canceled"
+ FinishReasonError FinishReason = "error"
+ FinishReasonPermissionDenied FinishReason = "permission_denied"
+
+ // Should never happen
+ FinishReasonUnknown FinishReason = "unknown"
+)
+
type ContentPart interface {
isPart()
}
@@ -73,13 +91,15 @@ type ToolResult struct {
ToolCallID string `json:"tool_call_id"`
Name string `json:"name"`
Content string `json:"content"`
+ Metadata string `json:"metadata"`
IsError bool `json:"is_error"`
}
func (ToolResult) isPart() {}
type Finish struct {
- Reason string `json:"reason"`
+ Reason FinishReason `json:"reason"`
+ Time int64 `json:"time"`
}
func (Finish) isPart() {}
@@ -89,6 +109,7 @@ type Message struct {
Role MessageRole
SessionID string
Parts []ContentPart
+ Model models.ModelID
CreatedAt int64
UpdatedAt int64
@@ -161,7 +182,16 @@ func (m *Message) IsFinished() bool {
return false
}
-func (m *Message) FinishReason() string {
+func (m *Message) FinishPart() *Finish {
+ for _, part := range m.Parts {
+ if c, ok := part.(Finish); ok {
+ return &c
+ }
+ }
+ return nil
+}
+
+func (m *Message) FinishReason() FinishReason {
for _, part := range m.Parts {
if c, ok := part.(Finish); ok {
return c.Reason
@@ -203,6 +233,40 @@ func (m *Message) AppendReasoningContent(delta string) {
}
}
+func (m *Message) FinishToolCall(toolCallID string) {
+ for i, part := range m.Parts {
+ if c, ok := part.(ToolCall); ok {
+ if c.ID == toolCallID {
+ m.Parts[i] = ToolCall{
+ ID: c.ID,
+ Name: c.Name,
+ Input: c.Input,
+ Type: c.Type,
+ Finished: true,
+ }
+ return
+ }
+ }
+ }
+}
+
+func (m *Message) AppendToolCallInput(toolCallID string, inputDelta string) {
+ for i, part := range m.Parts {
+ if c, ok := part.(ToolCall); ok {
+ if c.ID == toolCallID {
+ m.Parts[i] = ToolCall{
+ ID: c.ID,
+ Name: c.Name,
+ Input: c.Input + inputDelta,
+ Type: c.Type,
+ Finished: c.Finished,
+ }
+ return
+ }
+ }
+ }
+}
+
func (m *Message) AddToolCall(tc ToolCall) {
for i, part := range m.Parts {
if c, ok := part.(ToolCall); ok {
@@ -216,6 +280,15 @@ func (m *Message) AddToolCall(tc ToolCall) {
}
func (m *Message) SetToolCalls(tc []ToolCall) {
+ // remove any existing tool call part it could have multiple
+ parts := make([]ContentPart, 0)
+ for _, part := range m.Parts {
+ if _, ok := part.(ToolCall); ok {
+ continue
+ }
+ parts = append(parts, part)
+ }
+ m.Parts = parts
for _, toolCall := range tc {
m.Parts = append(m.Parts, toolCall)
}
@@ -231,8 +304,15 @@ func (m *Message) SetToolResults(tr []ToolResult) {
}
}
-func (m *Message) AddFinish(reason string) {
- m.Parts = append(m.Parts, Finish{Reason: reason})
+func (m *Message) AddFinish(reason FinishReason) {
+ // remove any existing finish part
+ for i, part := range m.Parts {
+ if _, ok := part.(Finish); ok {
+ m.Parts = slices.Delete(m.Parts, i, i+1)
+ break
+ }
+ }
+ m.Parts = append(m.Parts, Finish{Reason: reason, Time: time.Now().Unix()})
}
func (m *Message) AddImageURL(url, detail string) {
diff --git a/internal/message/message.go b/internal/message/message.go
index 13cf54048..20ace7b41 100644
--- a/internal/message/message.go
+++ b/internal/message/message.go
@@ -2,49 +2,51 @@ package message
import (
"context"
+ "database/sql"
"encoding/json"
"fmt"
+ "time"
"github.com/google/uuid"
- "github.com/kujtimiihoxha/termai/internal/db"
- "github.com/kujtimiihoxha/termai/internal/pubsub"
+ "github.com/kujtimiihoxha/opencode/internal/db"
+ "github.com/kujtimiihoxha/opencode/internal/llm/models"
+ "github.com/kujtimiihoxha/opencode/internal/pubsub"
)
type CreateMessageParams struct {
Role MessageRole
Parts []ContentPart
+ Model models.ModelID
}
type Service interface {
pubsub.Suscriber[Message]
- Create(sessionID string, params CreateMessageParams) (Message, error)
- Update(message Message) error
- Get(id string) (Message, error)
- List(sessionID string) ([]Message, error)
- Delete(id string) error
- DeleteSessionMessages(sessionID string) error
+ Create(ctx context.Context, sessionID string, params CreateMessageParams) (Message, error)
+ Update(ctx context.Context, message Message) error
+ Get(ctx context.Context, id string) (Message, error)
+ List(ctx context.Context, sessionID string) ([]Message, error)
+ Delete(ctx context.Context, id string) error
+ DeleteSessionMessages(ctx context.Context, sessionID string) error
}
type service struct {
*pubsub.Broker[Message]
- q db.Querier
- ctx context.Context
+ q db.Querier
}
-func NewService(ctx context.Context, q db.Querier) Service {
+func NewService(q db.Querier) Service {
return &service{
Broker: pubsub.NewBroker[Message](),
q: q,
- ctx: ctx,
}
}
-func (s *service) Delete(id string) error {
- message, err := s.Get(id)
+func (s *service) Delete(ctx context.Context, id string) error {
+ message, err := s.Get(ctx, id)
if err != nil {
return err
}
- err = s.q.DeleteMessage(s.ctx, message.ID)
+ err = s.q.DeleteMessage(ctx, message.ID)
if err != nil {
return err
}
@@ -52,7 +54,7 @@ func (s *service) Delete(id string) error {
return nil
}
-func (s *service) Create(sessionID string, params CreateMessageParams) (Message, error) {
+func (s *service) Create(ctx context.Context, sessionID string, params CreateMessageParams) (Message, error) {
if params.Role != Assistant {
params.Parts = append(params.Parts, Finish{
Reason: "stop",
@@ -63,11 +65,12 @@ func (s *service) Create(sessionID string, params CreateMessageParams) (Message,
return Message{}, err
}
- dbMessage, err := s.q.CreateMessage(s.ctx, db.CreateMessageParams{
+ dbMessage, err := s.q.CreateMessage(ctx, db.CreateMessageParams{
ID: uuid.New().String(),
SessionID: sessionID,
Role: string(params.Role),
Parts: string(partsJSON),
+ Model: sql.NullString{String: string(params.Model), Valid: true},
})
if err != nil {
return Message{}, err
@@ -80,14 +83,14 @@ func (s *service) Create(sessionID string, params CreateMessageParams) (Message,
return message, nil
}
-func (s *service) DeleteSessionMessages(sessionID string) error {
- messages, err := s.List(sessionID)
+func (s *service) DeleteSessionMessages(ctx context.Context, sessionID string) error {
+ messages, err := s.List(ctx, sessionID)
if err != nil {
return err
}
for _, message := range messages {
if message.SessionID == sessionID {
- err = s.Delete(message.ID)
+ err = s.Delete(ctx, message.ID)
if err != nil {
return err
}
@@ -96,32 +99,39 @@ func (s *service) DeleteSessionMessages(sessionID string) error {
return nil
}
-func (s *service) Update(message Message) error {
+func (s *service) Update(ctx context.Context, message Message) error {
parts, err := marshallParts(message.Parts)
if err != nil {
return err
}
- err = s.q.UpdateMessage(s.ctx, db.UpdateMessageParams{
- ID: message.ID,
- Parts: string(parts),
+ finishedAt := sql.NullInt64{}
+ if f := message.FinishPart(); f != nil {
+ finishedAt.Int64 = f.Time
+ finishedAt.Valid = true
+ }
+ err = s.q.UpdateMessage(ctx, db.UpdateMessageParams{
+ ID: message.ID,
+ Parts: string(parts),
+ FinishedAt: finishedAt,
})
if err != nil {
return err
}
+ message.UpdatedAt = time.Now().Unix()
s.Publish(pubsub.UpdatedEvent, message)
return nil
}
-func (s *service) Get(id string) (Message, error) {
- dbMessage, err := s.q.GetMessage(s.ctx, id)
+func (s *service) Get(ctx context.Context, id string) (Message, error) {
+ dbMessage, err := s.q.GetMessage(ctx, id)
if err != nil {
return Message{}, err
}
return s.fromDBItem(dbMessage)
}
-func (s *service) List(sessionID string) ([]Message, error) {
- dbMessages, err := s.q.ListMessagesBySession(s.ctx, sessionID)
+func (s *service) List(ctx context.Context, sessionID string) ([]Message, error) {
+ dbMessages, err := s.q.ListMessagesBySession(ctx, sessionID)
if err != nil {
return nil, err
}
@@ -145,6 +155,7 @@ func (s *service) fromDBItem(item db.Message) (Message, error) {
SessionID: item.SessionID,
Role: MessageRole(item.Role),
Parts: parts,
+ Model: models.ModelID(item.Model.String),
CreatedAt: item.CreatedAt,
UpdatedAt: item.UpdatedAt,
}, nil