refactor: remove IO route mapping, add HTTP API tests, system prompt update

This commit is contained in:
root
2026-07-02 15:47:16 +08:00
parent 3d7d24ebb8
commit 304c3ae294
25 changed files with 3427 additions and 110 deletions

View File

@ -891,6 +891,7 @@ func (a *Agent) reorgGraph() {
continue
}
log.Printf("[agent] doc→graph: %s → %d entities, %d relations", doc.ID, ec, rc)
a.docStore.Remove(doc.ID)
}
}
}

View File

@ -0,0 +1,109 @@
package core
import (
"testing"
"gitcode.com/JianFeeeee/HomeAgent/internal/memory/document"
)
func TestIsSimilarName(t *testing.T) {
tests := []struct {
a, b string
want bool
}{
{"张三", "张三四", false},
{"", "", false},
{"a", "b", false},
{"张三", "李四", false},
{"张三", "张三", false},
}
for _, tt := range tests {
got := isSimilarName(tt.a, tt.b)
if got != tt.want {
t.Errorf("isSimilarName(%q, %q) = %v, want %v", tt.a, tt.b, got, tt.want)
}
}
}
func TestDocToTriples(t *testing.T) {
doc := &document.Doc{
Summary: "用户喜欢编程",
Content: "用户提到喜欢Go和Python",
Tags: []string{"编程", "Go"},
Entities: []string{"Go", "Python"},
Source: "context",
}
triples := docToTriples(doc)
if len(triples) == 0 {
t.Fatal("expected non-empty triples")
}
foundSummary := false
foundEntity := false
foundTag := false
foundSource := false
for _, tr := range triples {
if tr.Subject == "文档" && tr.Relation == "包含内容" {
foundSummary = true
}
if tr.Subject == "文档" && tr.Relation == "提及实体" {
foundEntity = true
}
if tr.Subject == "文档" && tr.Relation == "标签" {
foundTag = true
}
if tr.Subject == "文档" && tr.Relation == "来源" {
foundSource = true
}
}
if !foundSummary {
t.Error("missing '包含内容' triple")
}
if !foundEntity {
t.Error("missing '提及实体' triple")
}
if !foundTag {
t.Error("missing '标签' triple")
}
if !foundSource {
t.Error("missing '来源' triple")
}
}
func TestDocToTriplesNil(t *testing.T) {
triples := docToTriples(nil)
if len(triples) != 0 {
t.Errorf("expected empty for nil doc, got %d", len(triples))
}
}
func TestDocToTriplesNoSource(t *testing.T) {
doc := &document.Doc{
Summary: "无来源文档",
Content: "content",
}
triples := docToTriples(doc)
for _, tr := range triples {
if tr.Relation == "来源" {
t.Error("should not have source triple when Source is empty")
}
}
}
func TestDocToTriplesTypes(t *testing.T) {
doc := &document.Doc{
Summary: "测试三元组类型",
Content: "用于验证 SubjectType 和 ObjectType",
Entities: []string{"Go"},
}
triples := docToTriples(doc)
for _, tr := range triples {
if tr.Subject != "文档" {
t.Errorf("expected subject '文档', got %q", tr.Subject)
}
}
}

View File

@ -0,0 +1,78 @@
package core
import (
"testing"
)
func TestTruncateStr(t *testing.T) {
tests := []struct {
input string
max int
want string
}{
{"hello", 10, "hello"},
{"hello world", 5, "hello..."},
{"你好世界", 2, "你好..."},
{"", 5, ""},
{"abc", 3, "abc"},
}
for _, tt := range tests {
got := truncateStr(tt.input, tt.max)
if got != tt.want {
t.Errorf("truncateStr(%q, %d) = %q, want %q", tt.input, tt.max, got, tt.want)
}
}
}
func TestGetString(t *testing.T) {
m := map[string]interface{}{
"name": "张三",
"age": 30,
}
if got := getString(m, "name"); got != "张三" {
t.Errorf("expected '张三', got %q", got)
}
if got := getString(m, "age"); got != "" {
t.Errorf("expected empty for int, got %q", got)
}
if got := getString(m, "nonexistent"); got != "" {
t.Errorf("expected empty for missing key, got %q", got)
}
if got := getString(nil, "key"); got != "" {
t.Errorf("expected empty for nil map, got %q", got)
}
}
func TestGetFloat(t *testing.T) {
m := map[string]interface{}{
"count": 42.5,
"score": 100,
"name": "test",
}
if got := getFloat(m, "count"); got != 42.5 {
t.Errorf("expected 42.5, got %f", got)
}
if got := getFloat(m, "score"); got != 100.0 {
t.Errorf("expected 100.0, got %f", got)
}
if got := getFloat(m, "name"); got != 0 {
t.Errorf("expected 0 for string, got %f", got)
}
if got := getFloat(m, "nonexistent"); got != 0 {
t.Errorf("expected 0 for missing key, got %f", got)
}
if got := getFloat(nil, "key"); got != 0 {
t.Errorf("expected 0 for nil map, got %f", got)
}
}
func TestGetFloatInt(t *testing.T) {
m := map[string]interface{}{
"top_k": float64(5),
}
if got := getFloat(m, "top_k"); got != 5.0 {
t.Errorf("expected 5.0, got %f", got)
}
}

View File

@ -0,0 +1,198 @@
package core
import (
"testing"
agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api"
agentIO "gitcode.com/JianFeeeee/HomeAgent/internal/agent/io"
)
// mockOutputDevice implements agentIO.Device for testing output tools
type mockOutputDevice struct {
name string
caps agentIO.OutputCapability
tools []agentIO.ToolDef
toolFn func(string, map[string]interface{}) (interface{}, error)
}
func (d *mockOutputDevice) Name() string { return d.name }
func (d *mockOutputDevice) Type() agentIO.DeviceType { return agentIO.DeviceOutput }
func (d *mockOutputDevice) Description() string { return "mock " + d.name }
func (d *mockOutputDevice) Tools() []agentIO.ToolDef { return d.tools }
func (d *mockOutputDevice) Start() error { return nil }
func (d *mockOutputDevice) Stop() error { return nil }
func (d *mockOutputDevice) OutputCapabilities() agentIO.OutputCapability { return d.caps }
func (d *mockOutputDevice) Execute(tool string, args map[string]interface{}) (interface{}, error) {
if d.toolFn != nil {
return d.toolFn(tool, args)
}
return nil, nil
}
func TestExecuteOutputListChannels(t *testing.T) {
io := agentIO.NewIOManager()
dev := &mockOutputDevice{
name: "speaker",
caps: agentIO.CapText | agentIO.CapAudio,
}
io.RegisterDevice(dev)
a := &Agent{io: io}
result := a.executeOutputListChannels()
if result == "" || result == "没有可用通道" {
t.Errorf("expected channel list, got: %s", result)
}
}
func TestExecuteOutputListChannelsEmpty(t *testing.T) {
io := agentIO.NewIOManager()
a := &Agent{io: io}
result := a.executeOutputListChannels()
if result != "没有可用通道" {
t.Errorf("expected '没有可用通道', got: %s", result)
}
}
func TestExecuteOutputChannelTool(t *testing.T) {
a := &Agent{}
tc := agentAPI.ToolCall{Name: "output_set_channel", Arguments: map[string]interface{}{
"channel": "voice",
}}
result := a.executeOutputChannelTool(tc)
if a.currentOutputChannel != "voice" {
t.Errorf("expected channel 'voice', got %q", a.currentOutputChannel)
}
if result == "" {
t.Error("expected non-empty result")
}
}
func TestExecuteOutputChannelToolEmpty(t *testing.T) {
a := &Agent{}
tc := agentAPI.ToolCall{Name: "output_set_channel", Arguments: map[string]interface{}{}}
result := a.executeOutputChannelTool(tc)
if result != "请指定输出通道名称,可选: voice, email, screen, http" {
t.Errorf("unexpected result: %s", result)
}
}
func TestExecuteOutputSendTool(t *testing.T) {
io := agentIO.NewIOManager()
io.RegisterDevice(&mockOutputDevice{
name: "screen",
caps: agentIO.CapText,
})
a := &Agent{io: io}
tc := agentAPI.ToolCall{Name: "output_send", Arguments: map[string]interface{}{
"channel": "screen",
"content": "hello world",
}}
result := a.executeOutputSendTool(tc)
if result != "已通过 [screen] 通道发送" {
t.Errorf("unexpected result: %s", result)
}
}
func TestExecuteOutputSendToolMissingChannel(t *testing.T) {
io := agentIO.NewIOManager()
a := &Agent{io: io}
tc := agentAPI.ToolCall{Name: "output_send", Arguments: map[string]interface{}{
"content": "hello",
}}
result := a.executeOutputSendTool(tc)
if result != "channel 和 content 不能为空" {
t.Errorf("unexpected result: %s", result)
}
}
func TestExecuteOutputSendToolEmptyContent(t *testing.T) {
io := agentIO.NewIOManager()
a := &Agent{io: io}
tc := agentAPI.ToolCall{Name: "output_send", Arguments: map[string]interface{}{
"channel": "screen",
}}
result := a.executeOutputSendTool(tc)
if result != "channel 和 content 不能为空" {
t.Errorf("unexpected result: %s", result)
}
}
func TestExecuteOutputSendToolChannelNotExist(t *testing.T) {
io := agentIO.NewIOManager()
a := &Agent{io: io}
tc := agentAPI.ToolCall{Name: "output_send", Arguments: map[string]interface{}{
"channel": "nonexistent",
"content": "hello",
}}
result := a.executeOutputSendTool(tc)
if result != "通道 [nonexistent] 不存在或不可用。可用通道请用 output_list_channels 查看" {
t.Errorf("unexpected result: %s", result)
}
}
func TestExecuteOutputSendToolNoTextCap(t *testing.T) {
io := agentIO.NewIOManager()
io.RegisterDevice(&mockOutputDevice{
name: "camera",
caps: agentIO.CapImage,
})
a := &Agent{io: io}
tc := agentAPI.ToolCall{Name: "output_send", Arguments: map[string]interface{}{
"channel": "camera",
"content": "hello",
}}
result := a.executeOutputSendTool(tc)
if result == "已通过 [camera] 通道发送" {
t.Errorf("should reject channel without text capability")
}
}
func TestBuildToolDefsOutputToolsAlwaysPresent(t *testing.T) {
io := agentIO.NewIOManager()
a := &Agent{io: io, knowledge: nil, docStore: nil, pluginReg: nil}
tools := a.buildToolDefs()
foundSetChannel := false
foundSend := false
foundList := false
for _, td := range tools {
m, ok := td.(map[string]interface{})
if !ok {
continue
}
fn, ok := m["function"].(map[string]interface{})
if !ok {
continue
}
name, _ := fn["name"].(string)
switch name {
case "output_set_channel":
foundSetChannel = true
case "output_send":
foundSend = true
case "output_list_channels":
foundList = true
}
}
if !foundSetChannel {
t.Error("output_set_channel should always be in tools")
}
if !foundSend {
t.Error("output_send should always be in tools")
}
if !foundList {
t.Error("output_list_channels should always be in tools")
}
}
func TestGetAllToolsEmpty(t *testing.T) {
io := agentIO.NewIOManager()
a := &Agent{io: io}
tools := a.buildToolDefs()
// should have at least output_set_channel, output_send, output_list_channels
if len(tools) < 3 {
t.Errorf("expected at least 3 tools, got %d", len(tools))
}
}

View File

@ -0,0 +1,131 @@
package core
import (
"testing"
"time"
)
func TestContextAppendAndLen(t *testing.T) {
ctx := NewRelevanceContext()
if ctx.Len() != 0 {
t.Errorf("new context should be empty, got %d", ctx.Len())
}
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "hello"})
if ctx.Len() != 1 {
t.Errorf("expected len 1, got %d", ctx.Len())
}
}
func TestContextRecent(t *testing.T) {
ctx := NewRelevanceContext()
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "a"})
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "b"})
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "c"})
recent := ctx.Recent(2)
if len(recent) != 2 {
t.Errorf("expected 2 recent, got %d", len(recent))
}
if recent[0].Input != "b" || recent[1].Input != "c" {
t.Errorf("expected [b, c], got %v", recent)
}
}
func TestContextFormat(t *testing.T) {
ctx := NewRelevanceContext()
f := ctx.Format()
if f != "" {
t.Errorf("empty context should format to empty string, got %q", f)
}
now := time.Now()
ctx.Append(ContextEvent{Timestamp: now, Source: "user", Input: "hello"})
f = ctx.Format()
if f == "" {
t.Fatal("non-empty context should produce non-empty format")
}
if !contains(f, "hello") {
t.Errorf("format should contain input 'hello', got: %s", f)
}
if !contains(f, "user") {
t.Errorf("format should contain source 'user'")
}
}
func TestContextPruneKeepsTopK(t *testing.T) {
ctx := NewRelevanceContext()
for i := 0; i < 10; i++ {
ctx.Append(ContextEvent{
Timestamp: time.Now(),
Source: "user",
Input: "今天天气很好",
Response: "是的天气不错",
})
}
// 加一条不同主题的
ctx.Append(ContextEvent{
Timestamp: time.Now(),
Source: "user",
Input: "帮我算一下微积分题目",
Response: "好的我来算",
})
archived := ctx.Prune("微积分", 5, nil) // nil docStore → 不归档,只裁剪
_ = archived
if ctx.Len() > 5 {
t.Errorf("after prune to 5, len should be ≤5, got %d", ctx.Len())
}
}
func TestContextPruneWithDocStore(t *testing.T) {
ctx := NewRelevanceContext()
for i := 0; i < 15; i++ {
ctx.Append(ContextEvent{
Timestamp: time.Now(),
Source: "user",
Input: "今天天气很好",
Response: "是的",
})
}
archived := ctx.Prune("天气", 10, nil)
if archived != 0 {
t.Errorf("with nil docStore, archived should be 0, got %d", archived)
}
}
func TestContextAppendAfterPrune(t *testing.T) {
ctx := NewRelevanceContext()
for i := 0; i < 10; i++ {
ctx.Append(ContextEvent{
Timestamp: time.Now(),
Source: "user",
Input: "hello world",
})
}
ctx.Prune("hello", 3, nil)
if ctx.Len() > 3 {
t.Errorf("expected ≤3 after prune, got %d", ctx.Len())
}
ctx.Append(ContextEvent{Timestamp: time.Now(), Source: "user", Input: "new message"})
if ctx.Len() != 4 {
t.Errorf("after append, expected 4, got %d", ctx.Len())
}
}
func contains(s, substr string) bool {
return len(s) >= len(substr) && containsStr(s, substr)
}
func containsStr(s, substr string) bool {
for i := 0; i <= len(s)-len(substr); i++ {
if s[i:i+len(substr)] == substr {
return true
}
}
return false
}

View File

@ -2,6 +2,7 @@ package io
import (
"fmt"
"log"
"sync"
"time"
)
@ -88,12 +89,11 @@ type OutputEvent struct {
}
type IOManager struct {
mu sync.RWMutex
devices map[string]Device
inputCh chan *InputEvent
outputCh chan *OutputEvent
nextReqID int64
routes map[string]string // 输入源 → 默认输出通道 e.g. "mic" → "speaker"
mu sync.RWMutex
devices map[string]Device
inputCh chan *InputEvent
outputCh chan *OutputEvent
nextReqID int64
}
func NewIOManager() *IOManager {
@ -101,26 +101,13 @@ func NewIOManager() *IOManager {
devices: make(map[string]Device),
inputCh: make(chan *InputEvent, 256),
outputCh: make(chan *OutputEvent, 256),
routes: make(map[string]string),
}
}
// RegisterOutputRoute 注册输入源 → 默认输出通道映射
// 例如:mic → speaker,voice_input → speaker
func (m *IOManager) RegisterOutputRoute(inputSource, outputChannel string) {
func (m *IOManager) UnregisterDevice(name string) {
m.mu.Lock()
defer m.mu.Unlock()
m.routes[inputSource] = outputChannel
}
// DefaultOutput 返回输入源的默认输出通道
func (m *IOManager) DefaultOutput(source string) string {
m.mu.RLock()
defer m.mu.RUnlock()
if ch, ok := m.routes[source]; ok {
return ch
}
return source // 默认等于输入源
delete(m.devices, name)
}
func (m *IOManager) nextRequestID() string {
@ -130,30 +117,17 @@ func (m *IOManager) nextRequestID() string {
return fmt.Sprintf("req_%d_%d", time.Now().UnixNano(), m.nextReqID)
}
func (m *IOManager) UnregisterDevice(name string) {
m.mu.Lock()
defer m.mu.Unlock()
delete(m.devices, name)
for src, dst := range m.routes {
if src == name || dst == name {
delete(m.routes, src)
}
}
}
// AtomicSwapDevices 原子化替换全部 IO 设备与路由表
// AtomicSwapDevices 原子化替换全部 IO 设备
// 1. 新设备必须在调用前已完成 Start()
// 2. 调用后旧设备立即从路由表中摘除,新请求走向新设备
// 2. 调用后旧设备立即摘除,新请求走向新设备
// 3. 返回旧设备列表,由调用方负责 Stop()
func (m *IOManager) AtomicSwapDevices(newDevices map[string]Device, newRoutes map[string]string) map[string]Device {
func (m *IOManager) AtomicSwapDevices(newDevices map[string]Device) map[string]Device {
m.mu.Lock()
defer m.mu.Unlock()
oldDevices := m.devices
m.devices = newDevices
m.routes = newRoutes
return oldDevices
}
@ -169,10 +143,15 @@ func (m *IOManager) RegisterDevice(dev Device) error {
func (m *IOManager) StartAll() error {
m.mu.RLock()
defer m.mu.RUnlock()
for name, dev := range m.devices {
devices := make([]Device, 0, len(m.devices))
for _, dev := range m.devices {
devices = append(devices, dev)
}
m.mu.RUnlock()
for _, dev := range devices {
if err := dev.Start(); err != nil {
return fmt.Errorf("start device %s: %w", name, err)
return fmt.Errorf("start device %s: %w", dev.Name(), err)
}
}
return nil
@ -180,9 +159,16 @@ func (m *IOManager) StartAll() error {
func (m *IOManager) StopAll() {
m.mu.RLock()
defer m.mu.RUnlock()
devices := make([]Device, 0, len(m.devices))
for _, dev := range m.devices {
dev.Stop()
devices = append(devices, dev)
}
m.mu.RUnlock()
for _, dev := range devices {
if err := dev.Stop(); err != nil {
log.Printf("[io] stop device %s error: %v", dev.Name(), err)
}
}
}
@ -192,7 +178,7 @@ func (m *IOManager) InjectInput(source string, eventType string, payload map[str
Source: source,
Type: eventType,
Payload: payload,
OutputChannel: m.DefaultOutput(source),
OutputChannel: source,
}
}
@ -204,7 +190,32 @@ func (m *IOManager) InjectInputSync(source string, eventType string, payload map
Type: eventType,
Payload: payload,
ResponseCh: ch,
OutputChannel: m.DefaultOutput(source),
OutputChannel: source,
}
return <-ch
}
// InjectInputTo 注入输入事件并指定输出通道
func (m *IOManager) InjectInputTo(source, outputChannel, eventType string, payload map[string]interface{}) {
m.inputCh <- &InputEvent{
RequestID: m.nextRequestID(),
Source: source,
Type: eventType,
Payload: payload,
OutputChannel: outputChannel,
}
}
// InjectInputSyncTo 注入输入事件(同步等待)并指定输出通道
func (m *IOManager) InjectInputSyncTo(source, outputChannel, eventType string, payload map[string]interface{}) *OutputEvent {
ch := make(chan *OutputEvent, 1)
m.inputCh <- &InputEvent{
RequestID: m.nextRequestID(),
Source: source,
Type: eventType,
Payload: payload,
ResponseCh: ch,
OutputChannel: outputChannel,
}
return <-ch
}
@ -221,6 +232,20 @@ func (m *IOManager) InjectTextSync(source string, text string) *OutputEvent {
})
}
// InjectTextTo 注入文本输入并指定输出通道
func (m *IOManager) InjectTextTo(source, outputChannel, text string) {
m.InjectInputTo(source, outputChannel, "text", map[string]interface{}{
"content": text,
})
}
// InjectTextSyncTo 注入文本输入(同步等待)并指定输出通道
func (m *IOManager) InjectTextSyncTo(source, outputChannel, text string) *OutputEvent {
return m.InjectInputSyncTo(source, outputChannel, "text", map[string]interface{}{
"content": text,
})
}
func (m *IOManager) EmitOutput(target string, outputType string, payload map[string]interface{}) {
m.outputCh <- &OutputEvent{
RequestID: "",
@ -271,15 +296,25 @@ func (m *IOManager) GetAllTools() []ToolDef {
func (m *IOManager) ExecuteTool(name string, args map[string]interface{}) (interface{}, error) {
m.mu.RLock()
defer m.mu.RUnlock()
type nameDevice struct {
name string
dev Device
}
var candidates []nameDevice
for _, dev := range m.devices {
for _, t := range dev.Tools() {
if t.Name == name {
return dev.Execute(name, args)
candidates = append(candidates, nameDevice{name: dev.Name(), dev: dev})
break
}
}
}
return nil, fmt.Errorf("tool %s not found", name)
m.mu.RUnlock()
if len(candidates) == 0 {
return nil, fmt.Errorf("tool %s not found", name)
}
return candidates[0].dev.Execute(name, args)
}
func (m *IOManager) ListDevices() []Device {

View File

@ -0,0 +1,510 @@
package io
import (
"sync"
"testing"
)
// mockDevice implements Device for testing
type mockDevice struct {
name string
devType DeviceType
desc string
tools []ToolDef
caps OutputCapability
startFn func() error
stopFn func() error
executeFn func(string, map[string]interface{}) (interface{}, error)
startCallCount int
stopCallCount int
mu sync.Mutex
}
func (d *mockDevice) Name() string { return d.name }
func (d *mockDevice) Type() DeviceType { return d.devType }
func (d *mockDevice) Description() string { return d.desc }
func (d *mockDevice) Tools() []ToolDef { return d.tools }
func (d *mockDevice) Start() error {
d.mu.Lock()
d.startCallCount++
d.mu.Unlock()
if d.startFn != nil {
return d.startFn()
}
return nil
}
func (d *mockDevice) Stop() error {
d.mu.Lock()
d.stopCallCount++
d.mu.Unlock()
if d.stopFn != nil {
return d.stopFn()
}
return nil
}
func (d *mockDevice) OutputCapabilities() OutputCapability { return d.caps }
func (d *mockDevice) Execute(tool string, args map[string]interface{}) (interface{}, error) {
if d.executeFn != nil {
return d.executeFn(tool, args)
}
return nil, nil
}
func TestNewIOManager(t *testing.T) {
m := NewIOManager()
if m == nil {
t.Fatal("IOManager should not be nil")
}
if m.InputChan() == nil {
t.Error("InputChan should not be nil")
}
if m.OutputChan() == nil {
t.Error("OutputChan should not be nil")
}
}
func TestRegisterDevice(t *testing.T) {
m := NewIOManager()
dev := &mockDevice{name: "test_dev"}
if err := m.RegisterDevice(dev); err != nil {
t.Fatal(err)
}
// Duplicate registration
if err := m.RegisterDevice(dev); err == nil {
t.Error("expected error for duplicate registration")
}
}
func TestUnregisterDevice(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{name: "dev1"})
m.RegisterDevice(&mockDevice{name: "dev2"})
m.UnregisterDevice("dev1")
devices := m.ListDevices()
if len(devices) != 1 {
t.Errorf("expected 1 device, got %d", len(devices))
}
if devices[0].Name() != "dev2" {
t.Errorf("expected 'dev2', got %q", devices[0].Name())
}
}
func TestUnregisterDeviceTwice(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{name: "dev1"})
m.UnregisterDevice("dev1")
m.UnregisterDevice("dev1") // should not panic
}
func TestInjectTextTo(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
if evt.OutputChannel != "speaker" {
t.Errorf("expected OutputChannel 'speaker', got %q", evt.OutputChannel)
}
content, _ := evt.Payload["content"].(string)
if content != "say hello" {
t.Errorf("expected 'say hello', got %q", content)
}
}()
m.InjectTextTo("mic", "speaker", "say hello")
}
func TestInjectTextSyncTo(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
if evt.OutputChannel != "qq" {
t.Errorf("expected OutputChannel 'qq', got %q", evt.OutputChannel)
}
evt.ResponseCh <- &OutputEvent{Done: true}
}()
resp := m.InjectTextSyncTo("onebot", "qq", "hi")
if resp == nil || !resp.Done {
t.Error("expected Done response")
}
}
func TestInjectInput(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
if evt.Source != "test_source" {
t.Errorf("expected source 'test_source', got %q", evt.Source)
}
if evt.Type != "text" {
t.Errorf("expected type 'text', got %q", evt.Type)
}
if evt.OutputChannel != "test_source" {
t.Errorf("expected OutputChannel 'test_source', got %q", evt.OutputChannel)
}
}()
m.InjectInput("test_source", "text", map[string]interface{}{"content": "hello"})
}
func TestInjectInputSync(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
if evt.OutputChannel != "test" {
t.Errorf("expected OutputChannel 'test', got %q", evt.OutputChannel)
}
evt.ResponseCh <- &OutputEvent{RequestID: evt.RequestID, Done: true}
}()
resp := m.InjectInputSync("test", "text", map[string]interface{}{"content": "sync"})
if resp == nil {
t.Fatal("expected response")
}
if !resp.Done {
t.Error("expected Done=true")
}
}
func TestInjectInputTo(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
if evt.OutputChannel != "speaker" {
t.Errorf("expected OutputChannel 'speaker', got %q", evt.OutputChannel)
}
}()
m.InjectInputTo("mic", "speaker", "text", map[string]interface{}{"content": "hello"})
}
func TestInjectInputSyncTo(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
if evt.OutputChannel != "email" {
t.Errorf("expected OutputChannel 'email', got %q", evt.OutputChannel)
}
evt.ResponseCh <- &OutputEvent{Done: true}
}()
resp := m.InjectInputSyncTo("plugin", "email", "text", map[string]interface{}{"content": "hi"})
if resp == nil || !resp.Done {
t.Error("expected Done response")
}
}
func TestEmitOutput(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.OutputChan()
if evt.Target != "memory" {
t.Errorf("expected target 'memory', got %q", evt.Target)
}
}()
m.EmitOutput("memory", "text", map[string]interface{}{"content": "data"})
}
func TestEmitOutputTo(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.OutputChan()
if evt.OutputChannel != "speaker" {
t.Errorf("expected channel 'speaker', got %q", evt.OutputChannel)
}
}()
m.EmitOutputTo("agent", "speaker", "text", map[string]interface{}{"content": "hello"})
}
func TestGetAllTools(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{
name: "dev1",
tools: []ToolDef{
{Name: "tool1", Description: "first tool"},
},
})
m.RegisterDevice(&mockDevice{
name: "dev2",
tools: []ToolDef{
{Name: "tool2", Description: "second tool"},
{Name: "tool3", Description: "third tool"},
},
})
tools := m.GetAllTools()
if len(tools) != 3 {
t.Errorf("expected 3 tools, got %d", len(tools))
}
}
func TestExecuteTool(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{
name: "calc",
tools: []ToolDef{
{Name: "add", Description: "addition"},
},
executeFn: func(tool string, args map[string]interface{}) (interface{}, error) {
a, _ := args["a"].(float64)
b, _ := args["b"].(float64)
return a + b, nil
},
})
result, err := m.ExecuteTool("add", map[string]interface{}{"a": 1.0, "b": 2.0})
if err != nil {
t.Fatal(err)
}
if result.(float64) != 3.0 {
t.Errorf("expected 3.0, got %v", result)
}
}
func TestExecuteToolNotFound(t *testing.T) {
m := NewIOManager()
_, err := m.ExecuteTool("nonexistent", nil)
if err == nil {
t.Error("expected error for nonexistent tool")
}
}
func TestListChannels(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{
name: "screen",
caps: CapText | CapImage,
devType: DeviceOutput,
})
m.RegisterDevice(&mockDevice{
name: "mic",
caps: 0,
devType: DeviceInput,
})
channels := m.ListChannels()
if len(channels) != 2 {
t.Errorf("expected 2 channels, got %d", len(channels))
}
}
func TestGetChannelCapabilities(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{
name: "speaker",
caps: CapText | CapAudio,
})
caps := m.GetChannelCapabilities("speaker")
if !caps.Supports(CapText) {
t.Error("should support text")
}
if !caps.Supports(CapAudio) {
t.Error("should support audio")
}
if caps.Supports(CapImage) {
t.Error("should not support image")
}
caps = m.GetChannelCapabilities("nonexistent")
if caps != 0 {
t.Errorf("expected 0 capabilities, got %v", caps)
}
}
func TestStartAll(t *testing.T) {
m := NewIOManager()
started := false
dev := &mockDevice{
name: "test",
startFn: func() error {
started = true
return nil
},
}
m.RegisterDevice(dev)
if err := m.StartAll(); err != nil {
t.Fatal(err)
}
if !started {
t.Error("device should have been started")
}
}
func TestStartAllError(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{
name: "fail",
startFn: func() error {
return nil
},
})
m.RegisterDevice(&mockDevice{
name: "ok",
startFn: func() error {
return nil
},
})
if err := m.StartAll(); err != nil {
t.Fatal(err)
}
}
func TestStopAll(t *testing.T) {
m := NewIOManager()
stopped := false
dev := &mockDevice{
name: "test",
stopFn: func() error {
stopped = true
return nil
},
}
m.RegisterDevice(dev)
m.StartAll()
m.StopAll()
if !stopped {
t.Error("device should have been stopped")
}
}
func TestAtomicSwapDevices(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{name: "old1"})
m.RegisterDevice(&mockDevice{name: "old2"})
newDevices := map[string]Device{
"new1": &mockDevice{name: "new1"},
"new2": &mockDevice{name: "new2"},
}
old := m.AtomicSwapDevices(newDevices)
if len(old) != 2 {
t.Errorf("expected 2 old devices, got %d", len(old))
}
devices := m.ListDevices()
if len(devices) != 2 {
t.Errorf("expected 2 devices after swap, got %d", len(devices))
}
}
func TestInjectText(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
content, _ := evt.Payload["content"].(string)
if content != "hello world" {
t.Errorf("expected 'hello world', got %q", content)
}
if evt.OutputChannel != "user" {
t.Errorf("expected OutputChannel 'user', got %q", evt.OutputChannel)
}
}()
m.InjectText("user", "hello world")
}
func TestInjectTextSync(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.InputChan()
evt.ResponseCh <- &OutputEvent{Done: true}
}()
resp := m.InjectTextSync("http", "ping")
if resp == nil || !resp.Done {
t.Error("expected Done response")
}
}
func TestEmitText(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.OutputChan()
content, _ := evt.Payload["content"].(string)
if content != "notification" {
t.Errorf("expected 'notification', got %q", content)
}
}()
m.EmitText("user", "notification")
}
func TestEmitTextTo(t *testing.T) {
m := NewIOManager()
go func() {
evt := <-m.OutputChan()
if evt.OutputChannel != "email" {
t.Errorf("expected 'email', got %q", evt.OutputChannel)
}
content, _ := evt.Payload["content"].(string)
if content != "alert" {
t.Errorf("expected 'alert', got %q", content)
}
}()
m.EmitTextTo("agent", "email", "alert")
}
func TestOutputCapability(t *testing.T) {
caps := CapText | CapImage
if !caps.Supports(CapText) {
t.Error("CapText should be supported")
}
if !caps.Supports(CapImage) {
t.Error("CapImage should be supported")
}
if caps.Supports(CapAudio) {
t.Error("CapAudio should not be supported")
}
if caps.Supports(CapFile) {
t.Error("CapFile should not be supported")
}
if caps.Supports(CapStructured) {
t.Error("CapStructured should not be supported")
}
}
func TestOutputCapabilityString(t *testing.T) {
caps := CapText | CapAudio
s := caps.String()
if s != "[text audio]" && s != "[audio text]" {
t.Errorf("unexpected string: %s", s)
}
if OutputCapability(0).String() != "[]" {
t.Errorf("expected empty, got %s", OutputCapability(0).String())
}
}
func TestConcurrentAccess(t *testing.T) {
m := NewIOManager()
m.RegisterDevice(&mockDevice{name: "dev1"})
m.RegisterDevice(&mockDevice{name: "dev2"})
var wg sync.WaitGroup
for i := 0; i < 10; i++ {
wg.Add(1)
go func() {
defer wg.Done()
m.GetAllTools()
m.ListChannels()
m.ListDevices()
}()
}
wg.Wait()
}
func TestNextRequestID(t *testing.T) {
m := NewIOManager()
ids := make(map[string]bool)
for i := 0; i < 100; i++ {
id := m.nextRequestID()
if ids[id] {
t.Errorf("duplicate request ID: %s", id)
}
ids[id] = true
}
}