- 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
163 lines
3.8 KiB
Go
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"
|
|
}
|