iop/apps/edge/internal/input/a2a/task_store.go
toki 672b4cdb5d refactor(bridge): 실행 경계 선행 안정화를 반영한다
원격 터미널 브리지 POC 전에 Client HTTP lifecycle과 Edge run result surface 계약을 고정해야 하므로 관련 구현, 테스트, 로드맵 상태를 함께 정리한다.
2026-06-07 08:40:36 +09:00

164 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 RunResult in a background goroutine and updates the stored Task.
func (s *TaskStore) collectBackground(taskID string, handle edgeservice.RunResult) {
go func() {
defer handle.Close()
s.drain(taskID, handle)
}()
}
// collectBlocking drains a RunResult synchronously and returns the final Task snapshot.
func (s *TaskStore) collectBlocking(taskID string, handle edgeservice.RunResult) *Task {
defer handle.Close()
return s.drain(taskID, handle)
}
func (s *TaskStore) drain(taskID string, handle edgeservice.RunResult) *Task {
var text strings.Builder
timeout := time.NewTimer(handle.WaitTimeout())
defer timeout.Stop()
stream := handle.Stream()
for {
select {
case nodeEvent := <-stream.NodeEvents:
if edgeservice.IsNodeDisconnected(nodeEvent) {
return s.setFailed(taskID, "node disconnected")
}
case event := <-stream.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"
}