diff --git a/backend/pkg/providers/providers.go b/backend/pkg/providers/providers.go
index a9a69c7d..b5d78463 100644
--- a/backend/pkg/providers/providers.go
+++ b/backend/pkg/providers/providers.go
@@ -5,6 +5,7 @@ import (
"crypto/rand"
"database/sql"
"encoding/json"
+ "errors"
"fmt"
"math"
"math/big"
@@ -317,7 +318,7 @@ func (pc *providerController) NewFlowProvider(
return nil, fmt.Errorf("failed to get primary docker image template: %w", err)
}
- image, err := prv.Call(ctx, pconfig.OptionsTypeSimple, imageTmpl)
+ image, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, imageTmpl)
if err != nil {
return nil, fmt.Errorf("failed to select primary docker image via llm call: %w", err)
}
@@ -330,7 +331,7 @@ func (pc *providerController) NewFlowProvider(
return nil, fmt.Errorf("failed to get language template: %w", err)
}
- language, err := prv.Call(ctx, pconfig.OptionsTypeSimple, languageTmpl)
+ language, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, languageTmpl)
if err != nil {
return nil, fmt.Errorf("failed to get language: %w", err)
}
@@ -346,7 +347,7 @@ func (pc *providerController) NewFlowProvider(
return nil, fmt.Errorf("failed to get flow title template: %w", err)
}
- title, err := prv.Call(ctx, pconfig.OptionsTypeSimple, titleTmpl)
+ title, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, titleTmpl)
if err != nil {
return nil, fmt.Errorf("failed to get flow title: %w", err)
}
@@ -476,7 +477,7 @@ func (pc *providerController) NewAssistantProvider(
return nil, fmt.Errorf("failed to get language template: %w", err)
}
- language, err := prv.Call(ctx, pconfig.OptionsTypeSimple, languageTmpl)
+ language, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, languageTmpl)
if err != nil {
return nil, fmt.Errorf("failed to get language: %w", err)
}
@@ -492,7 +493,7 @@ func (pc *providerController) NewAssistantProvider(
return nil, fmt.Errorf("failed to get flow title template: %w", err)
}
- title, err := prv.Call(ctx, pconfig.OptionsTypeSimple, titleTmpl)
+ title, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, titleTmpl)
if err != nil {
return nil, fmt.Errorf("failed to get flow title: %w", err)
}
@@ -1079,3 +1080,44 @@ func newAtomicInt64(seed int64) *atomic.Int64 {
number.Store(seed)
return &number
}
+
+// callWithSetupRetries wraps a single-shot LLM prompt call used during flow/
+// assistant bootstrap (docker image, language, and title selection) with the
+// same short retry-with-backoff already used for the agent execution loop
+// (see performSimpleChain/callWithRetries), so one transient error from the
+// LLM gateway (e.g. a bad gateway from a litellm proxy) does not fail flow or
+// assistant creation outright.
+func callWithSetupRetries(
+ ctx context.Context,
+ prv provider.Provider,
+ opt pconfig.ProviderOptionsType,
+ prompt string,
+) (string, error) {
+ var (
+ result string
+ err error
+ )
+
+ for idx := 0; idx <= maxRetriesToCallSimpleChain; idx++ {
+ if idx == maxRetriesToCallSimpleChain {
+ return "", fmt.Errorf("failed to call llm after %d retries: %w", idx, err)
+ }
+
+ result, err = prv.Call(ctx, opt, prompt)
+ if err == nil {
+ return result, nil
+ }
+
+ if errors.Is(err, context.Canceled) {
+ return "", err
+ }
+
+ select {
+ case <-ctx.Done():
+ return "", ctx.Err()
+ case <-time.After(delayBetweenRetries):
+ }
+ }
+
+ return "", err
+}
diff --git a/backend/pkg/providers/providers_test.go b/backend/pkg/providers/providers_test.go
index f988363d..52491677 100644
--- a/backend/pkg/providers/providers_test.go
+++ b/backend/pkg/providers/providers_test.go
@@ -2,8 +2,11 @@ package providers
import (
"context"
+ "errors"
+ "fmt"
"path/filepath"
"testing"
+ "time"
"pentagi/pkg/config"
"pentagi/pkg/database"
@@ -224,3 +227,86 @@ func TestAgentConfigPricesMatchCatalog(t *testing.T) {
}
}
}
+
+// fakeCallProvider overrides only Call(); embedding provider.Provider means
+// every other interface method is unimplemented and would panic on use,
+// which is fine since callWithSetupRetries only ever calls Call().
+type fakeCallProvider struct {
+ provider.Provider
+ callCount int
+ failTimes int
+ err error
+ result string
+}
+
+func (f *fakeCallProvider) Call(ctx context.Context, opt pconfig.ProviderOptionsType, prompt string) (string, error) {
+ f.callCount++
+ if f.callCount <= f.failTimes {
+ return "", f.err
+ }
+ return f.result, nil
+}
+
+func TestCallWithSetupRetries_SucceedsImmediately(t *testing.T) {
+ prv := &fakeCallProvider{result: "kali-linux"}
+
+ got, err := callWithSetupRetries(context.Background(), prv, pconfig.OptionsTypeSimple, "prompt")
+
+ require.NoError(t, err)
+ assert.Equal(t, "kali-linux", got)
+ assert.Equal(t, 1, prv.callCount, "a first-try success must not retry")
+}
+
+func TestCallWithSetupRetries_RetriesOnTransientErrorThenSucceeds(t *testing.T) {
+ // Exercises the real backoff once (a single delayBetweenRetries wait), so this
+ // test genuinely proves a transient 5xx from the LLM gateway self-heals
+ // instead of failing flow/assistant creation outright.
+ prv := &fakeCallProvider{
+ failTimes: 1,
+ err: fmt.Errorf("API returned unexpected status code: 502: bad gateway"),
+ result: "kali-linux",
+ }
+
+ got, err := callWithSetupRetries(context.Background(), prv, pconfig.OptionsTypeSimple, "prompt")
+
+ require.NoError(t, err)
+ assert.Equal(t, "kali-linux", got)
+ assert.Equal(t, 2, prv.callCount, "must retry exactly once after the transient failure")
+}
+
+func TestCallWithSetupRetries_ContextCanceledDuringWait_ReturnsWithoutFullBackoff(t *testing.T) {
+ prv := &fakeCallProvider{
+ failTimes: maxRetriesToCallSimpleChain, // always fails within the retry budget
+ err: errors.New("connection refused"),
+ }
+
+ ctx, cancel := context.WithCancel(context.Background())
+ go func() {
+ time.Sleep(20 * time.Millisecond)
+ cancel()
+ }()
+
+ start := time.Now()
+ _, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, "prompt")
+ elapsed := time.Since(start)
+
+ require.Error(t, err)
+ assert.ErrorIs(t, err, context.Canceled)
+ assert.Less(t, elapsed, delayBetweenRetries, "canceling context mid-wait must abort immediately, not wait out the full backoff")
+}
+
+func TestCallWithSetupRetries_ContextAlreadyCanceled_StopsWithoutRetrying(t *testing.T) {
+ prv := &fakeCallProvider{
+ failTimes: maxRetriesToCallSimpleChain,
+ err: context.Canceled,
+ }
+
+ ctx, cancel := context.WithCancel(context.Background())
+ cancel()
+
+ _, err := callWithSetupRetries(ctx, prv, pconfig.OptionsTypeSimple, "prompt")
+
+ require.Error(t, err)
+ assert.ErrorIs(t, err, context.Canceled)
+ assert.Equal(t, 1, prv.callCount, "a context.Canceled error from Call must stop retrying immediately")
+}
diff --git a/backend/pkg/templates/prompts/enricher.tmpl b/backend/pkg/templates/prompts/enricher.tmpl
index adceff6e..f7633391 100644
--- a/backend/pkg/templates/prompts/enricher.tmpl
+++ b/backend/pkg/templates/prompts/enricher.tmpl
@@ -171,7 +171,7 @@ Follow this prioritized approach to gather SUPPLEMENTARY information:
Search knowledge graph for episodic memory and execution history
Find what agents discovered and executed during operations
-recent_context, episode_context, successful_tools, entity_relationships
+recent_context, episode_context, successful_tools, entity_relationships (only after another search type returns a node UUID; needs center_node_uuid)
{{- end}}
diff --git a/backend/pkg/templates/prompts/memorist.tmpl b/backend/pkg/templates/prompts/memorist.tmpl
index abe7602f..18993ee6 100644
--- a/backend/pkg/templates/prompts/memorist.tmpl
+++ b/backend/pkg/templates/prompts/memorist.tmpl
@@ -47,6 +47,12 @@ ALWAYS search Graphiti BEFORE searching vector database:
- When asked about entities → Understand their relationships
+
+node_labels (PascalCase singular, use verbatim): Host, Port, Service, WebApp, Endpoint, Account, Vulnerability, Misconfiguration, Capability, Credential, ValidAccess, PrivChange, Tool, ToolExecution, Artifact, Evidence, Attempt, AttackTechnique.
+edge_types (UPPER_SNAKE_CASE, use verbatim): HAS_PORT, RUNS_SERVICE, HOSTS_APP, HAS_ENDPOINT, DETECTED_VULNERABILITY (scanner hit, unverified) → CONFIRMED_VULNERABILITY (validated) → HAS_VULNERABILITY (exploited), HAS_MISCONFIGURATION, AUTHENTICATES_TO, YIELDED_ACCESS, ESCALATED_VIA, PIVOTED_TO, ATTEMPTED_ON.
+Never invent a label/edge outside this list; if unsure, omit node_labels/edge_types and rely on the free-text query instead.
+
+
Choose the appropriate search type based on the information need:
@@ -70,15 +76,15 @@ Choose the appropriate search type based on the information need:
- When: Looking for working command examples, successful approaches
- Example: `search_type: "successful_tools", query: "successful nmap scans revealing services", min_mentions: 2`
-5. **entity_relationships** - Explore entity connections (requires entity UUID from prior search)
+5. **entity_relationships** - Explore entity connections (requires center_node_uuid from a prior search result; node_labels/edge_types are optional filters from the taxonomy reference above)
- Use: "What is connected to [entity]?"
- When: Understanding relationships between discovered entities
- Example: `search_type: "entity_relationships", query: "related vulnerabilities and services", center_node_uuid: "[uuid]", max_depth: 2`
-6. **entity_by_label** - Type-specific inventory (requires specific labels from prior discovery)
+6. **entity_by_label** - Type-specific inventory using node_labels from the taxonomy reference above
- Use: "List all [entity type] discovered"
- When: Creating inventories, generating comprehensive reports
- - Example: `search_type: "entity_by_label", query: "all discovered vulnerabilities", node_labels: ["VULNERABILITY"]`
+ - Example: `search_type: "entity_by_label", query: "all discovered vulnerabilities", node_labels: ["Vulnerability"]`
7. **diverse_results** - Get varied perspectives and alternatives
- Use: "What are different approaches/findings about [topic]?"
diff --git a/backend/pkg/templates/prompts/pentester.tmpl b/backend/pkg/templates/prompts/pentester.tmpl
index 685c3248..7404fc11 100644
--- a/backend/pkg/templates/prompts/pentester.tmpl
+++ b/backend/pkg/templates/prompts/pentester.tmpl
@@ -69,6 +69,12 @@ ALWAYS search Graphiti BEFORE attempting any significant action:
- After discovering entities → Understand their relationships
+
+node_labels (PascalCase singular, use verbatim): Host, Port, Service, WebApp, Endpoint, Account, Vulnerability, Misconfiguration, Capability, Credential, ValidAccess, PrivChange, Tool, ToolExecution, Artifact, Evidence, Attempt, AttackTechnique.
+edge_types (UPPER_SNAKE_CASE, use verbatim): HAS_PORT, RUNS_SERVICE, HOSTS_APP, HAS_ENDPOINT, DETECTED_VULNERABILITY (scanner hit, unverified) → CONFIRMED_VULNERABILITY (validated) → HAS_VULNERABILITY (exploited), HAS_MISCONFIGURATION, AUTHENTICATES_TO, YIELDED_ACCESS, ESCALATED_VIA, PIVOTED_TO, ATTEMPTED_ON.
+Never invent a label/edge outside this list; if unsure, omit node_labels/edge_types and rely on the free-text query instead.
+
+
Choose the appropriate search type based on your need:
@@ -87,7 +93,7 @@ Choose the appropriate search type based on your need:
- When: Need detailed context, understanding decision-making
- Example: `search_type: "episode_context", query: "pentester agent analysis of SSH vulnerability"`
-4. **entity_relationships** - Explore entity connections (can only be used after discovering an entity using other search types)
+4. **entity_relationships** - Explore entity connections (requires center_node_uuid from a prior search result; node_labels/edge_types are optional filters from the taxonomy reference above)
- Use: "What services/vulnerabilities are related to [entity]?"
- When: Investigating a specific IP, service, or vulnerability
- Example: `search_type: "entity_relationships", query: "services and vulnerabilities", center_node_uuid: "[uuid]", max_depth: 2`
@@ -97,10 +103,10 @@ Choose the appropriate search type based on your need:
- When: Current approach failing, need alternatives
- Example: `search_type: "diverse_results", query: "privilege escalation techniques on Linux", diversity_level: "high"`
-6. **entity_by_label** - Type-specific inventory (can only be used after discovering an entity using other search types, never use generic Entity label — only use specific labels like TechnicalFinding, Tool, AttackTechnique, etc.)
+6. **entity_by_label** - Type-specific inventory using node_labels from the taxonomy reference above — never a generic "Entity" label
- Use: "List all [entity type] we've discovered"
- When: Building inventories, generating reports
- - Example: `search_type: "entity_by_label", query: "all discovered vulnerabilities", node_labels: ["VULNERABILITY"]`
+ - Example: `search_type: "entity_by_label", query: "all discovered vulnerabilities", node_labels: ["Vulnerability"]`
diff --git a/backend/pkg/tools/args.go b/backend/pkg/tools/args.go
index ed018716..6618d679 100644
--- a/backend/pkg/tools/args.go
+++ b/backend/pkg/tools/args.go
@@ -159,15 +159,15 @@ type WebSearchAction struct {
}
type GraphitiSearchAction struct {
- SearchType String `json:"search_type" jsonschema:"required,type=string,enum=temporal_window,enum=entity_relationships,enum=diverse_results,enum=episode_context,enum=successful_tools,enum=recent_context,enum=entity_by_label" jsonschema_description:"Type of search to perform: temporal_window (time-bounded search), entity_relationships (graph traversal from an entity), diverse_results (anti-redundancy search), episode_context (full agent reasoning and tool outputs), successful_tools (proven techniques), recent_context (latest findings), entity_by_label (type-specific entity search)"`
+ SearchType String `json:"search_type" jsonschema:"required,type=string,enum=temporal_window,enum=entity_relationships,enum=diverse_results,enum=episode_context,enum=successful_tools,enum=recent_context,enum=entity_by_label" jsonschema_description:"Type of search to perform: temporal_window (time-bounded search), entity_relationships (graph traversal from an entity), diverse_results (anti-redundancy search), episode_context (full agent reasoning and tool outputs), successful_tools (proven techniques), recent_context (latest findings), entity_by_label (type-specific entity search, REQUIRES node_labels)"`
Query string `json:"query" jsonschema:"required" jsonschema_description:"Technical-channel payload — natural language query against the team's temporal knowledge graph. ALWAYS written in English regardless of the engagement language: the graph is indexed in English and shared across all engagements; non-English queries will fail to retrieve stored episodic memory."`
MaxResults *Int64 `json:"max_results,omitempty" jsonschema:"title=Maximum Results,type=integer" jsonschema_description:"Maximum number of results to return (default varies by search type)"`
TimeStart string `json:"time_start,omitempty" jsonschema_description:"Start of time window (ISO 8601 format, required for temporal_window)"`
TimeEnd string `json:"time_end,omitempty" jsonschema_description:"End of time window (ISO 8601 format, required for temporal_window)"`
CenterNodeUUID string `json:"center_node_uuid,omitempty" jsonschema_description:"UUID of entity to search from (required for entity_relationships)"`
MaxDepth *Int64 `json:"max_depth,omitempty" jsonschema:"title=Maximum Depth,type=integer" jsonschema_description:"Maximum graph traversal depth (default: 2, max: 3, for entity_relationships)"`
- NodeLabels []string `json:"node_labels,omitempty" jsonschema_description:"Filter to specific node types (e.g., ['IP_ADDRESS', 'SERVICE', 'VULNERABILITY'])"`
- EdgeTypes []string `json:"edge_types,omitempty" jsonschema_description:"Filter to specific relationship types (e.g., ['HAS_PORT', 'EXPLOITS'])"`
+ NodeLabels []string `json:"node_labels,omitempty" jsonschema_description:"Filter to specific node types — EXACT taxonomy names, PascalCase singular (e.g., ['Host', 'Service', 'Vulnerability']). REQUIRED (non-empty) when search_type is entity_by_label; optional filter otherwise"`
+ EdgeTypes []string `json:"edge_types,omitempty" jsonschema_description:"Filter to specific relationship types — EXACT taxonomy names, UPPER_SNAKE_CASE (e.g., ['HAS_PORT', 'HAS_VULNERABILITY'])"`
DiversityLevel String `json:"diversity_level,omitempty" jsonschema:"type=string,enum=low,enum=medium,enum=high" jsonschema_description:"How much diversity to prioritize (default: medium, for diverse_results)"`
MinMentions *Int64 `json:"min_mentions,omitempty" jsonschema:"title=Minimum Mentions,type=integer" jsonschema_description:"Minimum episode mentions (default: 2, for successful_tools)"`
RecencyWindow String `json:"recency_window,omitempty" jsonschema:"type=string,enum=1h,enum=6h,enum=24h,enum=7d" jsonschema_description:"How far back to search (default: 24h, for recent_context)"`
diff --git a/backend/pkg/tools/browser.go b/backend/pkg/tools/browser.go
index 3bc368b5..4f5f7f31 100644
--- a/backend/pkg/tools/browser.go
+++ b/backend/pkg/tools/browser.go
@@ -22,9 +22,10 @@ import (
)
const (
- minMdContentSize = 50
- minHtmlContentSize = 300
- minImgContentSize = 2048
+ minMdContentSize = 50
+ minHtmlContentSize = 300
+ minImgContentSize = 2048
+ maxScraperErrorBodyBytes = 512
)
// nonHTMLExtensions lists URL path suffixes that point to resources the scraper
@@ -511,6 +512,17 @@ func (b *browser) callScraper(url string) ([]byte, error) {
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
+ if resp.StatusCode >= 500 {
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, maxScraperErrorBodyBytes+1))
+ if preview := strings.TrimSpace(string(body)); preview != "" {
+ if truncated := len(body) > maxScraperErrorBodyBytes; truncated {
+ preview = preview[:maxScraperErrorBodyBytes] + "... [truncated]"
+ }
+ return nil, fmt.Errorf(
+ "unexpected resp code for scraper '%s': %d, response: %s", url, resp.StatusCode, preview,
+ )
+ }
+ }
return nil, fmt.Errorf("unexpected resp code for scraper '%s': %d", url, resp.StatusCode)
}
diff --git a/backend/pkg/tools/browser_test.go b/backend/pkg/tools/browser_test.go
index 0075ff90..92da6b42 100644
--- a/backend/pkg/tools/browser_test.go
+++ b/backend/pkg/tools/browser_test.go
@@ -413,6 +413,73 @@ func TestGetHTML_EmptyContent_ReturnsError(t *testing.T) {
}
}
+func TestCallScraper_ServerError5xx_IncludesBodyPreview(t *testing.T) {
+ ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.WriteHeader(http.StatusBadGateway)
+ fmt.Fprint(w, "502 Bad Gateway")
+ }))
+ defer ts.Close()
+
+ b := &browser{flowID: 1}
+
+ _, err := b.callScraper(ts.URL)
+ if err == nil {
+ t.Fatal("callScraper() should error on 5xx response")
+ }
+ if !strings.Contains(err.Error(), "502") {
+ t.Errorf("callScraper() error should mention the status code, got: %v", err)
+ }
+ if !strings.Contains(err.Error(), "502 Bad Gateway") {
+ t.Errorf("callScraper() error should include the response body, got: %v", err)
+ }
+}
+
+func TestCallScraper_ServerError5xx_BodyTruncatedAt1024Bytes(t *testing.T) {
+ hugeBody := strings.Repeat("x", 2000)
+
+ ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.WriteHeader(http.StatusInternalServerError)
+ fmt.Fprint(w, hugeBody)
+ }))
+ defer ts.Close()
+
+ b := &browser{flowID: 1}
+
+ _, err := b.callScraper(ts.URL)
+ if err == nil {
+ t.Fatal("callScraper() should error on 5xx response")
+ }
+ if strings.Count(err.Error(), "x") >= 2000 {
+ t.Fatalf("expected the 2000-byte body to be truncated to the 1024-byte cap, got error of length %d", len(err.Error()))
+ }
+ if !strings.Contains(err.Error(), "truncated") {
+ t.Errorf("callScraper() error should indicate truncation, got: %v", err)
+ }
+}
+
+func TestCallScraper_ClientError4xx_NoBodyPreview(t *testing.T) {
+ ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.WriteHeader(http.StatusNotFound)
+ fmt.Fprint(w, "not found body")
+ }))
+ defer ts.Close()
+
+ b := &browser{flowID: 1}
+
+ _, err := b.callScraper(ts.URL)
+ if err == nil {
+ t.Fatal("callScraper() should error on 4xx response")
+ }
+ if !strings.Contains(err.Error(), "404") {
+ t.Errorf("callScraper() error should mention the status code, got: %v", err)
+ }
+ // 4xx means our own request was malformed, not the scraper backend failing,
+ // so the body preview (only meaningful for 5xx) must not be included.
+ if strings.Contains(err.Error(), "not found body") {
+ t.Errorf("callScraper() error should not include the response body for 4xx, got: %v", err)
+ }
+}
+
func TestGetHTML_BinaryURL_ReturnsError(t *testing.T) {
b := &browser{flowID: 1, scPubURL: "http://127.0.0.1:1"}
diff --git a/backend/pkg/tools/graphiti_search.go b/backend/pkg/tools/graphiti_search.go
index 99f048e2..8dc10d9c 100644
--- a/backend/pkg/tools/graphiti_search.go
+++ b/backend/pkg/tools/graphiti_search.go
@@ -15,6 +15,7 @@ import (
obs "pentagi/pkg/observability"
"pentagi/pkg/observability/langfuse"
+ "github.com/google/uuid"
"github.com/sirupsen/logrus"
)
@@ -361,6 +362,12 @@ func (t *graphitiSearchTool) handleEntityRelationshipsSearch(
if args.CenterNodeUUID == "" {
return "", fmt.Errorf("center_node_uuid is required for entity_relationships search")
}
+ if _, err := uuid.Parse(args.CenterNodeUUID); err != nil {
+ return "", fmt.Errorf(
+ "center_node_uuid must be a valid UUID copied verbatim from the 'UUID:' field of a prior "+
+ "graphiti_search result, got %q", args.CenterNodeUUID,
+ )
+ }
maxResults := args.MaxResults.Int()
if maxResults <= 0 {
@@ -544,7 +551,10 @@ func (t *graphitiSearchTool) handleEntityByLabelSearch(
observationObject *graphiti.Observation,
) (string, error) {
if len(args.NodeLabels) == 0 {
- return "", fmt.Errorf("node_labels is required for entity_by_label search")
+ return "", fmt.Errorf(
+ "node_labels is required for entity_by_label search: pass one or more EXACT taxonomy node names " +
+ `(PascalCase singular), e.g. node_labels: ["Host", "Service", "Vulnerability"]`,
+ )
}
maxResults := args.MaxResults.Int()
diff --git a/backend/pkg/tools/graphiti_search_test.go b/backend/pkg/tools/graphiti_search_test.go
index 9b47b8a6..4001ac34 100644
--- a/backend/pkg/tools/graphiti_search_test.go
+++ b/backend/pkg/tools/graphiti_search_test.go
@@ -29,7 +29,10 @@ func (s *stubGraphitiSearcher) TemporalWindowSearch(
func (s *stubGraphitiSearcher) EntityRelationshipsSearch(
ctx context.Context, req graphiti.EntityRelationshipSearchRequest,
) (*graphiti.EntityRelationshipSearchResponse, error) {
- return nil, s.err
+ if s.err != nil {
+ return nil, s.err
+ }
+ return &graphiti.EntityRelationshipSearchResponse{}, nil
}
func (s *stubGraphitiSearcher) DiverseResultsSearch(
@@ -59,7 +62,10 @@ func (s *stubGraphitiSearcher) RecentContextSearch(
func (s *stubGraphitiSearcher) EntityByLabelSearch(
ctx context.Context, req graphiti.EntityByLabelSearchRequest,
) (*graphiti.EntityByLabelSearchResponse, error) {
- return nil, s.err
+ if s.err != nil {
+ return nil, s.err
+ }
+ return &graphiti.EntityByLabelSearchResponse{}, nil
}
// fakeNetError mimics the *url.Error shape produced by http.Client.Do on a
@@ -175,6 +181,76 @@ func TestGraphitiSearchTool_Handle_ValidationError_StaysHard(t *testing.T) {
}
}
+func TestGraphitiSearchTool_Handle_EntityByLabel_MissingNodeLabels_GivesActionableError(t *testing.T) {
+ tool := NewGraphitiSearchTool(1, nil, nil, &stubGraphitiSearcher{enabled: true})
+
+ args := []byte(`{"search_type":"entity_by_label","query":"test query","message":"m"}`)
+ _, err := tool.Handle(t.Context(), GraphitiSearchToolName, args)
+
+ if err == nil {
+ t.Fatal("expected a hard failure when node_labels is missing for entity_by_label, got nil error")
+ }
+ if !strings.Contains(err.Error(), "node_labels is required") {
+ t.Fatalf("expected error to state node_labels is required, got: %v", err)
+ }
+ if !strings.Contains(err.Error(), "Vulnerability") {
+ t.Fatalf("expected error to include a real taxonomy example value to guide the LLM, got: %v", err)
+ }
+}
+
+func TestGraphitiSearchTool_Handle_EntityByLabel_WithNodeLabels_Succeeds(t *testing.T) {
+ tool := NewGraphitiSearchTool(1, nil, nil, &stubGraphitiSearcher{enabled: true})
+
+ args := []byte(`{"search_type":"entity_by_label","query":"test query","node_labels":["Vulnerability"],"message":"m"}`)
+ _, err := tool.Handle(t.Context(), GraphitiSearchToolName, args)
+
+ if err != nil {
+ t.Fatalf("expected no error when node_labels is present, got: %v", err)
+ }
+}
+
+func TestGraphitiSearchTool_Handle_EntityRelationships_MissingCenterNodeUUID_StaysHard(t *testing.T) {
+ tool := NewGraphitiSearchTool(1, nil, nil, &stubGraphitiSearcher{enabled: true})
+
+ args := []byte(`{"search_type":"entity_relationships","query":"test query","message":"m"}`)
+ _, err := tool.Handle(t.Context(), GraphitiSearchToolName, args)
+
+ if err == nil || !strings.Contains(err.Error(), "center_node_uuid is required") {
+ t.Fatalf("expected hard 'center_node_uuid is required' error, got: %v", err)
+ }
+}
+
+func TestGraphitiSearchTool_Handle_EntityRelationships_MalformedCenterNodeUUID_GivesActionableError(t *testing.T) {
+ tool := NewGraphitiSearchTool(1, nil, nil, &stubGraphitiSearcher{enabled: true})
+
+ // Simulates a diagnostic string, truncated ID, or otherwise non-UUID value
+ // ending up in center_node_uuid (e.g. a hallucinated or mangled value) -
+ // this must be rejected before ever reaching the graph backend.
+ args := []byte(`{"search_type":"entity_relationships","query":"test query","center_node_uuid":"not-a-real-uuid","message":"m"}`)
+ _, err := tool.Handle(t.Context(), GraphitiSearchToolName, args)
+
+ if err == nil {
+ t.Fatal("expected a hard failure for a malformed center_node_uuid, got nil error")
+ }
+ if !strings.Contains(err.Error(), "must be a valid UUID") {
+ t.Fatalf("expected error to explain the UUID requirement, got: %v", err)
+ }
+ if !strings.Contains(err.Error(), "not-a-real-uuid") {
+ t.Fatalf("expected error to echo back the offending value, got: %v", err)
+ }
+}
+
+func TestGraphitiSearchTool_Handle_EntityRelationships_ValidCenterNodeUUID_Succeeds(t *testing.T) {
+ tool := NewGraphitiSearchTool(1, nil, nil, &stubGraphitiSearcher{enabled: true})
+
+ args := []byte(`{"search_type":"entity_relationships","query":"test query","center_node_uuid":"f7b95dfc-ee58-4a8b-8d85-582cf117b4df","message":"m"}`)
+ _, err := tool.Handle(t.Context(), GraphitiSearchToolName, args)
+
+ if err != nil {
+ t.Fatalf("expected no error for a well-formed center_node_uuid, got: %v", err)
+ }
+}
+
func TestGraphitiSearchTool_Handle_InvalidRecencyWindow_StaysHard(t *testing.T) {
tool := NewGraphitiSearchTool(1, nil, nil, &stubGraphitiSearcher{enabled: true})