opencode/internal/message/message.go

283 lines
6.4 KiB
Go
Raw Normal View History

2025-03-24 05:25:31 +08:00
package message
import (
"context"
2025-04-12 08:01:45 +08:00
"database/sql"
2025-03-24 05:25:31 +08:00
"encoding/json"
2025-04-03 21:20:15 +08:00
"fmt"
2025-04-19 22:35:45 +08:00
"time"
2025-03-24 05:25:31 +08:00
"github.com/google/uuid"
2025-04-25 00:25:52 +08:00
"github.com/opencode-ai/opencode/internal/db"
"github.com/opencode-ai/opencode/internal/llm/models"
"github.com/opencode-ai/opencode/internal/pubsub"
2025-03-24 05:25:31 +08:00
)
2025-03-28 05:35:48 +08:00
type CreateMessageParams struct {
2025-04-03 21:20:15 +08:00
Role MessageRole
Parts []ContentPart
2025-04-12 08:01:45 +08:00
Model models.ModelID
2025-03-24 05:25:31 +08:00
}
type Service interface {
pubsub.Suscriber[Message]
2025-04-13 19:17:17 +08:00
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
2025-03-24 05:25:31 +08:00
}
type service struct {
*pubsub.Broker[Message]
2025-04-13 19:17:17 +08:00
q db.Querier
2025-03-24 05:25:31 +08:00
}
2025-04-13 19:17:17 +08:00
func NewService(q db.Querier) Service {
2025-04-03 21:20:15 +08:00
return &service{
Broker: pubsub.NewBroker[Message](),
q: q,
}
}
2025-04-13 19:17:17 +08:00
func (s *service) Delete(ctx context.Context, id string) error {
message, err := s.Get(ctx, id)
2025-03-28 05:35:48 +08:00
if err != nil {
return err
}
2025-04-13 19:17:17 +08:00
err = s.q.DeleteMessage(ctx, message.ID)
2025-03-28 05:35:48 +08:00
if err != nil {
return err
}
s.Publish(pubsub.DeletedEvent, message)
return nil
}
2025-04-13 19:17:17 +08:00
func (s *service) Create(ctx context.Context, sessionID string, params CreateMessageParams) (Message, error) {
2025-04-03 21:20:15 +08:00
if params.Role != Assistant {
params.Parts = append(params.Parts, Finish{
Reason: "stop",
})
2025-03-28 05:35:48 +08:00
}
2025-04-03 21:20:15 +08:00
partsJSON, err := marshallParts(params.Parts)
2025-03-24 05:25:31 +08:00
if err != nil {
return Message{}, err
}
2025-04-03 21:20:15 +08:00
2025-04-13 19:17:17 +08:00
dbMessage, err := s.q.CreateMessage(ctx, db.CreateMessageParams{
2025-04-03 21:20:15 +08:00
ID: uuid.New().String(),
SessionID: sessionID,
Role: string(params.Role),
Parts: string(partsJSON),
2025-04-12 08:01:45 +08:00
Model: sql.NullString{String: string(params.Model), Valid: true},
2025-03-24 05:25:31 +08:00
})
if err != nil {
return Message{}, err
}
2025-03-28 05:35:48 +08:00
message, err := s.fromDBItem(dbMessage)
2025-03-24 05:25:31 +08:00
if err != nil {
2025-03-28 05:35:48 +08:00
return Message{}, err
2025-03-24 05:25:31 +08:00
}
2025-03-28 05:35:48 +08:00
s.Publish(pubsub.CreatedEvent, message)
return message, nil
2025-03-24 05:25:31 +08:00
}
2025-04-13 19:17:17 +08:00
func (s *service) DeleteSessionMessages(ctx context.Context, sessionID string) error {
messages, err := s.List(ctx, sessionID)
2025-03-24 05:25:31 +08:00
if err != nil {
return err
}
for _, message := range messages {
if message.SessionID == sessionID {
2025-04-13 19:17:17 +08:00
err = s.Delete(ctx, message.ID)
2025-03-24 05:25:31 +08:00
if err != nil {
return err
}
}
}
return nil
}
2025-04-13 19:17:17 +08:00
func (s *service) Update(ctx context.Context, message Message) error {
2025-04-03 21:20:15 +08:00
parts, err := marshallParts(message.Parts)
2025-03-28 05:35:48 +08:00
if err != nil {
return err
}
2025-04-12 08:01:45 +08:00
finishedAt := sql.NullInt64{}
if f := message.FinishPart(); f != nil {
finishedAt.Int64 = f.Time
finishedAt.Valid = true
}
2025-04-13 19:17:17 +08:00
err = s.q.UpdateMessage(ctx, db.UpdateMessageParams{
2025-04-12 08:01:45 +08:00
ID: message.ID,
Parts: string(parts),
FinishedAt: finishedAt,
2025-03-28 05:35:48 +08:00
})
if err != nil {
return err
}
2025-04-19 22:35:45 +08:00
message.UpdatedAt = time.Now().Unix()
2025-03-28 05:35:48 +08:00
s.Publish(pubsub.UpdatedEvent, message)
return nil
}
2025-04-13 19:17:17 +08:00
func (s *service) Get(ctx context.Context, id string) (Message, error) {
dbMessage, err := s.q.GetMessage(ctx, id)
2025-03-24 05:25:31 +08:00
if err != nil {
return Message{}, err
}
2025-03-28 05:35:48 +08:00
return s.fromDBItem(dbMessage)
2025-03-24 05:25:31 +08:00
}
2025-04-13 19:17:17 +08:00
func (s *service) List(ctx context.Context, sessionID string) ([]Message, error) {
dbMessages, err := s.q.ListMessagesBySession(ctx, sessionID)
2025-03-24 05:25:31 +08:00
if err != nil {
return nil, err
}
messages := make([]Message, len(dbMessages))
for i, dbMessage := range dbMessages {
2025-03-28 05:35:48 +08:00
messages[i], err = s.fromDBItem(dbMessage)
if err != nil {
return nil, err
}
2025-03-24 05:25:31 +08:00
}
return messages, nil
}
2025-03-28 05:35:48 +08:00
func (s *service) fromDBItem(item db.Message) (Message, error) {
2025-04-03 21:20:15 +08:00
parts, err := unmarshallParts([]byte(item.Parts))
if err != nil {
return Message{}, err
2025-03-28 05:35:48 +08:00
}
2025-04-03 21:20:15 +08:00
return Message{
ID: item.ID,
SessionID: item.SessionID,
Role: MessageRole(item.Role),
Parts: parts,
Model: models.ModelID(item.Model.String),
2025-04-03 21:20:15 +08:00
CreatedAt: item.CreatedAt,
UpdatedAt: item.UpdatedAt,
}, nil
}
2025-03-28 05:35:48 +08:00
2025-04-03 21:20:15 +08:00
type partType string
const (
reasoningType partType = "reasoning"
textType partType = "text"
imageURLType partType = "image_url"
binaryType partType = "binary"
toolCallType partType = "tool_call"
toolResultType partType = "tool_result"
finishType partType = "finish"
)
type partWrapper struct {
Type partType `json:"type"`
Data ContentPart `json:"data"`
}
func marshallParts(parts []ContentPart) ([]byte, error) {
wrappedParts := make([]partWrapper, len(parts))
for i, part := range parts {
var typ partType
switch part.(type) {
case ReasoningContent:
typ = reasoningType
case TextContent:
typ = textType
case ImageURLContent:
typ = imageURLType
case BinaryContent:
typ = binaryType
case ToolCall:
typ = toolCallType
case ToolResult:
typ = toolResultType
case Finish:
typ = finishType
default:
return nil, fmt.Errorf("unknown part type: %T", part)
2025-03-28 05:35:48 +08:00
}
2025-04-03 21:20:15 +08:00
wrappedParts[i] = partWrapper{
Type: typ,
Data: part,
}
}
return json.Marshal(wrappedParts)
2025-03-24 05:25:31 +08:00
}
2025-04-03 21:20:15 +08:00
func unmarshallParts(data []byte) ([]ContentPart, error) {
temp := []json.RawMessage{}
if err := json.Unmarshal(data, &temp); err != nil {
return nil, err
2025-03-24 05:25:31 +08:00
}
2025-04-03 21:20:15 +08:00
parts := make([]ContentPart, 0)
for _, rawPart := range temp {
var wrapper struct {
Type partType `json:"type"`
Data json.RawMessage `json:"data"`
}
if err := json.Unmarshal(rawPart, &wrapper); err != nil {
return nil, err
}
switch wrapper.Type {
case reasoningType:
part := ReasoningContent{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
parts = append(parts, part)
case textType:
part := TextContent{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
parts = append(parts, part)
case imageURLType:
part := ImageURLContent{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
case binaryType:
part := BinaryContent{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
parts = append(parts, part)
case toolCallType:
part := ToolCall{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
parts = append(parts, part)
case toolResultType:
part := ToolResult{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
parts = append(parts, part)
case finishType:
part := Finish{}
if err := json.Unmarshal(wrapper.Data, &part); err != nil {
return nil, err
}
parts = append(parts, part)
default:
return nil, fmt.Errorf("unknown part type: %s", wrapper.Type)
}
}
return parts, nil
2025-03-24 05:25:31 +08:00
}