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() if stream.Events == nil { return s.setFailed(taskID, "run stream unavailable") } for { select { case nodeEvent, ok := <-stream.NodeEvents: if !ok { stream.NodeEvents = nil continue } if edgeservice.IsNodeDisconnected(nodeEvent) { return s.setFailed(taskID, "node disconnected") } case event, ok := <-stream.Events: if !ok { return s.setFailed(taskID, "run stream closed") } 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" }