mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-10-03 15:53:56 +00:00
refactor: remove IO route mapping, add HTTP API tests, system prompt update
This commit is contained in:
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
109
internal/agent/core/agent_functions_test.go
Normal file
109
internal/agent/core/agent_functions_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
78
internal/agent/core/agent_helpers_test.go
Normal file
78
internal/agent/core/agent_helpers_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
198
internal/agent/core/agent_tools_test.go
Normal file
198
internal/agent/core/agent_tools_test.go
Normal 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))
|
||||
}
|
||||
}
|
||||
131
internal/agent/core/context_test.go
Normal file
131
internal/agent/core/context_test.go
Normal 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
|
||||
}
|
||||
Reference in New Issue
Block a user