141 lines
5.4 KiB
Go
141 lines
5.4 KiB
Go
package ollama
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"github.com/example/glpi-ai-agent/internal/model"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestAnalyseStructured(t *testing.T) {
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var body map[string]any
|
|
_ = json.NewDecoder(r.Body).Decode(&body)
|
|
format, _ := body["format"].(map[string]any)
|
|
if format == nil {
|
|
t.Error("missing schema")
|
|
} else if props, _ := format["properties"].(map[string]any); props != nil {
|
|
if category, _ := props["category"].(map[string]any); category != nil {
|
|
if categoryProps, _ := category["properties"].(map[string]any); categoryProps != nil {
|
|
if _, exists := categoryProps["change"]; exists {
|
|
t.Error("category schema must not let the model decide change=true/false")
|
|
}
|
|
}
|
|
}
|
|
}
|
|
options, _ := body["options"].(map[string]any)
|
|
if options["num_predict"] != float64(256) {
|
|
t.Errorf("unexpected num_predict: %v", options["num_predict"])
|
|
}
|
|
if body["keep_alive"] != "10m0s" {
|
|
t.Errorf("unexpected keep_alive: %v", body["keep_alive"])
|
|
}
|
|
if body["think"] != false {
|
|
t.Errorf("unexpected think: %v", body["think"])
|
|
}
|
|
json.NewEncoder(w).Encode(map[string]any{"message": map[string]any{"content": `{"category":{"id":1,"confidence":0.9},"reply":{"allowed":false,"confidence":0.1,"knowledge_id":""},"reason":"ok"}`}})
|
|
}))
|
|
defer srv.Close()
|
|
c := New(srv.URL, "m", "e", "de-DE", "formal", time.Second, 256, 10*time.Minute, false, 1, 1)
|
|
d, err := c.Analyse(context.Background(), model.Ticket{ID: 1}, []model.Category{{ID: 1}}, nil, model.ContextSnapshot{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if d.Reason != "ok" {
|
|
t.Fatalf("unexpected %+v", d)
|
|
}
|
|
}
|
|
|
|
func TestAnalyseRetriesInvalidJSON(t *testing.T) {
|
|
calls := 0
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
calls++
|
|
content := `{"category":`
|
|
if calls > 1 {
|
|
content = `{"category":{"id":2,"confidence":0.95},"reply":{"allowed":false,"confidence":0.1,"knowledge_id":""},"reason":"ok"}`
|
|
}
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"message": map[string]any{"content": content}})
|
|
}))
|
|
defer srv.Close()
|
|
c := New(srv.URL, "m", "e", "de-DE", "formal", time.Second, 768, time.Minute, false, 1, 1)
|
|
d, err := c.Analyse(context.Background(), model.Ticket{ID: 1}, []model.Category{{ID: 2, Name: "Active Directory"}}, nil, model.ContextSnapshot{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if calls != 2 || d.Category.ID != 2 {
|
|
t.Fatalf("calls=%d decision=%+v", calls, d)
|
|
}
|
|
}
|
|
|
|
func TestEmbedDisablesSilentTruncation(t *testing.T) {
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if body["truncate"] != false {
|
|
t.Fatalf("truncate=%v, want false", body["truncate"])
|
|
}
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"embeddings": [][]float64{{1, 0}}})
|
|
}))
|
|
defer srv.Close()
|
|
c := New(srv.URL, "m", "e", "de-DE", "formal", time.Second, 256, time.Minute, false, 1, 0)
|
|
v, err := c.Embed(context.Background(), []string{"test"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if len(v) != 1 {
|
|
t.Fatalf("embeddings=%d", len(v))
|
|
}
|
|
}
|
|
|
|
func TestAnalyseRetriesAllowedReplyWithoutKnowledgeID(t *testing.T) {
|
|
calls := 0
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
calls++
|
|
content := `{"category":{"id":2,"confidence":0.95},"reply":{"allowed":true,"confidence":0.95,"knowledge_id":""},"reason":"passt"}`
|
|
if calls > 1 {
|
|
content = `{"category":{"id":2,"confidence":0.95},"reply":{"allowed":true,"confidence":0.95,"knowledge_id":"GLPI-KB-1"},"reason":"passt"}`
|
|
}
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"message": map[string]any{"content": content}})
|
|
}))
|
|
defer srv.Close()
|
|
c := New(srv.URL, "m", "e", "de-DE", "formal", time.Second, 768, time.Minute, false, 1, 1)
|
|
hits := []model.KnowledgeHit{{Doc: model.KnowledgeDoc{ID: "GLPI-KB-1", Title: "Benutzeranmeldung"}}}
|
|
d, err := c.Analyse(context.Background(), model.Ticket{ID: 11}, []model.Category{{ID: 2, Name: "Active Directory"}}, hits, model.ContextSnapshot{})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if calls != 2 || d.Reply.KnowledgeID != "GLPI-KB-1" {
|
|
t.Fatalf("calls=%d decision=%+v", calls, d)
|
|
}
|
|
}
|
|
|
|
func TestAnalyseDoesNotExposeRichAnswerHTMLToModel(t *testing.T) {
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var body map[string]any
|
|
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
messages, _ := body["messages"].([]any)
|
|
for _, raw := range messages {
|
|
m, _ := raw.(map[string]any)
|
|
content, _ := m["content"].(string)
|
|
if strings.Contains(content, "RICH_SECRET_MARKUP") {
|
|
t.Fatalf("rich answer HTML leaked into Ollama prompt: %s", content)
|
|
}
|
|
}
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"message": map[string]any{"content": `{"category":{"id":2,"confidence":0.95},"reply":{"allowed":false,"confidence":0.1,"knowledge_id":""},"reason":"ok"}`}})
|
|
}))
|
|
defer srv.Close()
|
|
c := New(srv.URL, "m", "e", "de-DE", "formal", time.Second, 768, time.Minute, false, 1, 0)
|
|
hits := []model.KnowledgeHit{{Doc: model.KnowledgeDoc{ID: "GLPI-KB-1", Title: "Login", Text: "plain", Answer: "plain", AnswerHTML: `<p>RICH_SECRET_MARKUP</p>`}}}
|
|
if _, err := c.Analyse(context.Background(), model.Ticket{ID: 1}, []model.Category{{ID: 2}}, hits, model.ContextSnapshot{}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|