Skip to content

Commit 43d79b7

Browse files
committed
RDBC-1059 Add AiAgentToolSubAgent and AiAgentConfiguration.sub_agents
1 parent 29b9bf2 commit 43d79b7

5 files changed

Lines changed: 174 additions & 1 deletion

File tree

ravendb/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -95,6 +95,7 @@
9595
AiAgentParameterValueType,
9696
AiAgentToolAction,
9797
AiAgentToolQuery,
98+
AiAgentToolSubAgent,
9899
AiAgentPersistenceConfiguration,
99100
AiAgentChatTrimmingConfiguration,
100101
AiAgentSummarizationByTokens,

ravendb/documents/operations/ai/agents/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
AiAgentTruncateChat,
1313
AiAgentHistoryConfiguration,
1414
)
15+
from .ai_agent_tool_sub_agent import AiAgentToolSubAgent
1516

1617
from .add_or_update_ai_agent_operation import (
1718
AddOrUpdateAiAgentOperation,
@@ -45,6 +46,7 @@
4546
"AiAgentToolAction",
4647
"AiAgentToolQuery",
4748
"AiAgentToolQueryOptions",
49+
"AiAgentToolSubAgent",
4850
"AiAgentPersistenceConfiguration",
4951
"AiAgentChatTrimmingConfiguration",
5052
"AiAgentSummarizationByTokens",

ravendb/documents/operations/ai/agents/ai_agent_configuration.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -306,7 +306,10 @@ def __init__(
306306
chat_trimming: AiAgentChatTrimmingConfiguration = None,
307307
max_model_iterations_per_call: int = None,
308308
disabled: bool = False,
309+
sub_agents: List["AiAgentToolSubAgent"] = None,
309310
):
311+
from ravendb.documents.operations.ai.agents.ai_agent_tool_sub_agent import AiAgentToolSubAgent
312+
310313
self.name = name
311314
self.connection_string_name = connection_string_name
312315
self.system_prompt = system_prompt
@@ -320,10 +323,10 @@ def __init__(
320323
self.chat_trimming: Optional[AiAgentChatTrimmingConfiguration] = chat_trimming
321324
self.max_model_iterations_per_call: Optional[int] = max_model_iterations_per_call
322325
self.disabled: bool = disabled
326+
self.sub_agents: List[AiAgentToolSubAgent] = sub_agents or []
323327

324328
@staticmethod
325329
def _normalize_parameters(parameters: List[Union[str, AiAgentParameter]]) -> List[AiAgentParameter]:
326-
"""Convert a list of strings or AiAgentParameter objects to a list of AiAgentParameter objects."""
327330
if not parameters:
328331
return []
329332
result = []
@@ -344,6 +347,7 @@ def to_json(self) -> Dict[str, Any]:
344347
"OutputSchema": self.output_schema,
345348
"Queries": [q.to_json() for q in self.queries],
346349
"Actions": [a.to_json() for a in self.actions],
350+
"SubAgents": [s.to_json() for s in self.sub_agents],
347351
"Persistence": self.persistence.to_json() if self.persistence else None,
348352
"Parameters": [p.to_json() for p in self.parameters],
349353
"ChatTrimming": self.chat_trimming.to_json() if self.chat_trimming else None,
@@ -369,6 +373,12 @@ def from_json(cls, json_dict: Dict[str, Any]) -> AiAgentConfiguration:
369373
if actions_data:
370374
instance.actions = [AiAgentToolAction.from_json(a) for a in actions_data]
371375

376+
from ravendb.documents.operations.ai.agents.ai_agent_tool_sub_agent import AiAgentToolSubAgent
377+
378+
sub_agents_data = json_dict.get("subAgents") or json_dict.get("SubAgents")
379+
if sub_agents_data:
380+
instance.sub_agents = [AiAgentToolSubAgent.from_json(s) for s in sub_agents_data]
381+
372382
persistence_data = json_dict.get("persistence") or json_dict.get("Persistence")
373383
if persistence_data:
374384
instance.persistence = AiAgentPersistenceConfiguration.from_json(persistence_data)
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
from __future__ import annotations
2+
from typing import Any, Dict, Optional
3+
4+
5+
class AiAgentToolSubAgent:
6+
"""Server-side sub-agent the model can call from within a parent agent run."""
7+
8+
def __init__(self, identifier: Optional[str] = None, description: Optional[str] = None):
9+
self.identifier = identifier
10+
self.description = description
11+
12+
def to_json(self) -> Dict[str, Any]:
13+
return {
14+
"Identifier": self.identifier,
15+
"Description": self.description,
16+
}
17+
18+
@classmethod
19+
def from_json(cls, json_dict: Dict[str, Any]) -> AiAgentToolSubAgent:
20+
return cls(
21+
identifier=json_dict.get("identifier") or json_dict.get("Identifier"),
22+
description=json_dict.get("description") or json_dict.get("Description"),
23+
)
Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,137 @@
1+
"""
2+
Integration tests against a live RavenDB 7.2.x server for the AI agent
3+
configuration fields added in 7.2.3:
4+
* AiAgentConfiguration.sub_agents (+ AiAgentToolSubAgent)
5+
* AiAgentParameter.policy (AiAgentParameterPolicy)
6+
* AiAgentParameter.type (AiAgentParameterValueType)
7+
8+
The license guard mirrors the existing AI agent tests in this directory.
9+
"""
10+
11+
import os
12+
import unittest
13+
14+
from ravendb import (
15+
AiAgentConfiguration,
16+
AiAgentParameter,
17+
AiAgentParameterPolicy,
18+
AiAgentParameterValueType,
19+
AiAgentToolSubAgent,
20+
)
21+
from ravendb.documents.operations.ai import AiConnectionString, AiModelType
22+
from ravendb.documents.operations.ai.agents import (
23+
AddOrUpdateAiAgentOperation,
24+
DeleteAiAgentOperation,
25+
)
26+
from ravendb.documents.operations.ai.open_ai_settings import OpenAiSettings
27+
from ravendb.documents.operations.connection_string.put_connection_string_operation import PutConnectionStringOperation
28+
from ravendb.tests.test_base import TestBase
29+
30+
31+
@unittest.skipIf(os.environ.get("RAVENDB_LICENSE") is None, "Insufficient license permissions. Skipping on CI/CD.")
32+
class TestAiAgentConfigExtensionsIntegration(TestBase):
33+
CONNECTION_STRING_NAME = "test-ai-agent-cs-ext"
34+
35+
def setUp(self):
36+
super().setUp()
37+
ai_connection_string = AiConnectionString(
38+
name=self.CONNECTION_STRING_NAME,
39+
identifier=self.CONNECTION_STRING_NAME,
40+
model_type=AiModelType.CHAT,
41+
openai_settings=OpenAiSettings(
42+
api_key="dummy-api-key",
43+
endpoint="https://api.openai.com/v1",
44+
model="gpt-4",
45+
),
46+
)
47+
self.store.maintenance.send(PutConnectionStringOperation(ai_connection_string))
48+
self._created_agent_ids = []
49+
50+
def tearDown(self):
51+
for agent_id in self._created_agent_ids:
52+
try:
53+
self.store.maintenance.send(DeleteAiAgentOperation(agent_id))
54+
except Exception:
55+
pass
56+
super().tearDown()
57+
58+
# ---- sub_agents ----
59+
60+
def test_sub_agents_round_trip_through_get_agent(self):
61+
agent = AiAgentConfiguration(
62+
name="ParentAgent",
63+
identifier="test-sub-agent-parent",
64+
connection_string_name=self.CONNECTION_STRING_NAME,
65+
system_prompt="Dispatch to sub-agents as needed.",
66+
sample_object='{"answer": "..."}',
67+
sub_agents=[
68+
AiAgentToolSubAgent(identifier="benefits-agent", description="Handles benefit questions"),
69+
AiAgentToolSubAgent(identifier="attendance-agent", description="Tracks PTO and attendance"),
70+
],
71+
)
72+
result = self.store.ai.add_or_update_agent(agent)
73+
self._created_agent_ids.append(result.identifier)
74+
75+
fetched = self.store.ai.get_agents(result.identifier).ai_agents[0]
76+
identifiers = sorted(s.identifier for s in fetched.sub_agents)
77+
self.assertEqual(["attendance-agent", "benefits-agent"], identifiers)
78+
79+
descriptions = {s.identifier: s.description for s in fetched.sub_agents}
80+
self.assertEqual("Handles benefit questions", descriptions["benefits-agent"])
81+
self.assertEqual("Tracks PTO and attendance", descriptions["attendance-agent"])
82+
83+
# ---- AiAgentParameter.policy ----
84+
85+
def test_parameter_policy_forbid_model_generation_round_trips(self):
86+
agent = AiAgentConfiguration(
87+
name="ParamPolicyAgent",
88+
identifier="test-param-policy",
89+
connection_string_name=self.CONNECTION_STRING_NAME,
90+
system_prompt="Test parameter policy.",
91+
sample_object='{"answer": "..."}',
92+
parameters=[
93+
AiAgentParameter(
94+
name="user_id",
95+
description="Hidden user id",
96+
send_to_model=False,
97+
policy=AiAgentParameterPolicy.FORBID_MODEL_GENERATION,
98+
),
99+
AiAgentParameter(name="country", description="The country to filter by."),
100+
],
101+
)
102+
result = self.store.ai.add_or_update_agent(agent)
103+
self._created_agent_ids.append(result.identifier)
104+
105+
fetched = self.store.ai.get_agents(result.identifier).ai_agents[0]
106+
by_name = {p.name: p for p in fetched.parameters}
107+
self.assertEqual(AiAgentParameterPolicy.FORBID_MODEL_GENERATION, by_name["user_id"].policy)
108+
# Default-valued parameter still comes back with the default policy.
109+
self.assertEqual(AiAgentParameterPolicy.DEFAULT, by_name["country"].policy)
110+
111+
# ---- AiAgentParameter.type ----
112+
113+
def test_parameter_value_type_round_trips(self):
114+
agent = AiAgentConfiguration(
115+
name="ParamTypeAgent",
116+
identifier="test-param-type",
117+
connection_string_name=self.CONNECTION_STRING_NAME,
118+
system_prompt="Test parameter types.",
119+
sample_object='{"answer": "..."}',
120+
parameters=[
121+
AiAgentParameter(name="email", description="Email", type=AiAgentParameterValueType.STRING),
122+
AiAgentParameter(name="age", description="Age", type=AiAgentParameterValueType.NUMBER),
123+
AiAgentParameter(name="tags", description="Tags", type=AiAgentParameterValueType.ARRAY_OF_STRING),
124+
],
125+
)
126+
result = self.store.ai.add_or_update_agent(agent)
127+
self._created_agent_ids.append(result.identifier)
128+
129+
fetched = self.store.ai.get_agents(result.identifier).ai_agents[0]
130+
by_name = {p.name: p for p in fetched.parameters}
131+
self.assertEqual(AiAgentParameterValueType.STRING, by_name["email"].type)
132+
self.assertEqual(AiAgentParameterValueType.NUMBER, by_name["age"].type)
133+
self.assertEqual(AiAgentParameterValueType.ARRAY_OF_STRING, by_name["tags"].type)
134+
135+
136+
if __name__ == "__main__":
137+
unittest.main()

0 commit comments

Comments
 (0)