iop/apps/edge/internal/input/a2a/task_store.go
toki b4124f0bd6 feat: CLI setup, edge/node transport refactor, and infrastructure updates
- Add CLI core setup for edge and node services
- Refactor edge transport layer (server, integration tests)
- Refactor node transport layer (parser, session, heartbeat, client)
- Add main_test.go files for edge and node commands
- Add input package for edge service
- Add go.work and go.work.sum for workspace support
- Update configs, docs, and project rules
2026-05-20 16:37:42 +09:00

163 lines
3.8 KiB
Go

package a2a
import (
"strings"
"sync"
"time"
edgeservice "iop/apps/edge/internal/service"
iop "iop/proto/gen/iop"
)
// TaskStore is an in-memory registry of A2A tasks indexed by task (run) ID.
type TaskStore struct {
mu sync.RWMutex
tasks map[string]*Task
}
func newTaskStore() *TaskStore {
return &TaskStore{tasks: make(map[string]*Task)}
}
func (s *TaskStore) get(id string) (*Task, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
t, ok := s.tasks[id]
if !ok {
return nil, false
}
return snapshotTask(t), true
}
// snapshotTask returns a deep copy of t so callers cannot observe future mutations.
// Must be called with s.mu held for reading (or from a context where t cannot be
// concurrently modified, e.g. already holding the write lock).
func snapshotTask(t *Task) *Task {
cp := Task{
ID: t.ID,
Status: TaskStatus{State: t.Status.State},
}
if t.Status.Message != nil {
msgCopy := copyMessage(*t.Status.Message)
cp.Status.Message = &msgCopy
}
if len(t.Artifacts) > 0 {
cp.Artifacts = make([]Artifact, len(t.Artifacts))
for i, a := range t.Artifacts {
cp.Artifacts[i] = Artifact{
Name: a.Name,
Parts: append([]Part(nil), a.Parts...),
}
}
}
if len(t.History) > 0 {
cp.History = make([]Message, len(t.History))
for i, m := range t.History {
cp.History[i] = copyMessage(m)
}
}
return &cp
}
func copyMessage(m Message) Message {
return Message{
Role: m.Role,
Parts: append([]Part(nil), m.Parts...),
}
}
// create stores a new working Task and returns the task ID (not the internal pointer).
func (s *TaskStore) create(id string, userMsg Message) string {
t := &Task{
ID: id,
Status: TaskStatus{State: StateWorking},
History: []Message{userMsg},
}
s.mu.Lock()
s.tasks[id] = t
s.mu.Unlock()
return id
}
// collectBackground drains a RunHandle in a background goroutine and updates the stored Task.
func (s *TaskStore) collectBackground(taskID string, handle *edgeservice.RunHandle) {
go func() {
defer handle.Close()
s.drain(taskID, handle)
}()
}
// collectBlocking drains a RunHandle synchronously and returns the final Task snapshot.
func (s *TaskStore) collectBlocking(taskID string, handle *edgeservice.RunHandle) *Task {
defer handle.Close()
return s.drain(taskID, handle)
}
func (s *TaskStore) drain(taskID string, handle *edgeservice.RunHandle) *Task {
var text strings.Builder
timeout := time.NewTimer(handle.WaitTimeout())
defer timeout.Stop()
for {
select {
case nodeEvent := <-handle.NodeEvents:
if edgeservice.IsNodeDisconnected(nodeEvent) {
return s.setFailed(taskID, "node disconnected")
}
case event := <-handle.Events:
if event == nil {
continue
}
switch event.GetType() {
case "delta":
text.WriteString(event.GetDelta())
case "complete":
artifact := Artifact{Parts: []Part{{Type: "text", Text: text.String()}}}
s.mu.Lock()
if t, ok := s.tasks[taskID]; ok {
t.Artifacts = []Artifact{artifact}
t.Status = TaskStatus{State: StateCompleted}
}
s.mu.Unlock()
t, _ := s.get(taskID)
return t
case "error":
msg := errorMsg(event)
return s.setFailed(taskID, msg)
case "cancelled":
s.mu.Lock()
if t, ok := s.tasks[taskID]; ok {
t.Status = TaskStatus{State: StateCanceled}
}
s.mu.Unlock()
t, _ := s.get(taskID)
return t
}
case <-timeout.C:
return s.setFailed(taskID, "run timed out")
}
}
}
func (s *TaskStore) setFailed(taskID, msg string) *Task {
s.mu.Lock()
if t, ok := s.tasks[taskID]; ok {
t.Status = TaskStatus{State: StateFailed, Message: &Message{
Role: "agent",
Parts: []Part{{Type: "text", Text: msg}},
}}
}
s.mu.Unlock()
t, _ := s.get(taskID)
return t
}
func errorMsg(event *iop.RunEvent) string {
if m := event.GetError(); m != "" {
return m
}
if m := event.GetMessage(); m != "" {
return m
}
return "run failed"
}