Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 74 additions & 0 deletions tests/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -558,6 +558,80 @@ def test_find_dfiq(self, mock_get):
result = self.api.find_dfiq(name="not_found", dfiq_type="scenario")
self.assertIsNone(result)

@patch("yeti.api.requests.Session.post")
def test_search_agent_personas(self, mock_post):
mock_response = MagicMock()
mock_response.content = b'{"personas": [{"name": "Default"}], "total": 1}'
mock_post.return_value = mock_response

result = self.api.search_agent_personas()
self.assertEqual(result, [{"name": "Default"}])
mock_post.assert_called_with(
"http://fake-url/api/v2/agentpersonas/search",
json={"name": "", "count": 50, "page": 0},
)

@patch("yeti.api.requests.Session.post")
def test_search_agent_personas_filtered(self, mock_post):
mock_response = MagicMock()
mock_response.content = b'{"personas": [], "total": 0}'
mock_post.return_value = mock_response

self.api.search_agent_personas(name="SOC", enabled=True, count=10, page=2)
mock_post.assert_called_with(
"http://fake-url/api/v2/agentpersonas/search",
json={"name": "SOC", "count": 10, "page": 2, "enabled": True},
)

@patch("yeti.api.requests.Session.post")
def test_search_agent_personas_enabled_false_is_sent(self, mock_post):
"""False must reach the API; only None means "either"."""
mock_response = MagicMock()
mock_response.content = b'{"personas": [], "total": 0}'
mock_post.return_value = mock_response

self.api.search_agent_personas(enabled=False)
mock_post.assert_called_with(
"http://fake-url/api/v2/agentpersonas/search",
json={"name": "", "count": 50, "page": 0, "enabled": False},
)

@patch("yeti.api.requests.Session.get")
def test_get_agent_persona(self, mock_get):
mock_response = MagicMock()
mock_response.content = b'{"id": "1", "name": "Default"}'
mock_get.return_value = mock_response

result = self.api.get_agent_persona("1")
self.assertEqual(result, {"id": "1", "name": "Default"})
mock_get.assert_called_with("http://fake-url/api/v2/agentpersonas/1")

@patch("yeti.api.requests.Session.post")
def test_new_agent_persona(self, mock_post):
mock_response = MagicMock()
mock_response.content = b'{"id": "new_persona"}'
mock_post.return_value = mock_response

result = self.api.new_agent_persona({"name": "Default"})
self.assertEqual(result, {"id": "new_persona"})
mock_post.assert_called_with(
"http://fake-url/api/v2/agentpersonas/",
json={"persona": {"name": "Default"}},
)

@patch("yeti.api.requests.Session.patch")
def test_patch_agent_persona(self, mock_patch):
mock_response = MagicMock()
mock_response.content = b'{"id": "patched_persona"}'
mock_patch.return_value = mock_response

result = self.api.patch_agent_persona("1", {"name": "Renamed"})
self.assertEqual(result, {"id": "patched_persona"})
mock_patch.assert_called_with(
"http://fake-url/api/v2/agentpersonas/1",
json={"persona": {"name": "Renamed"}},
)


if __name__ == "__main__":
unittest.main()
82 changes: 82 additions & 0 deletions yeti/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -982,3 +982,85 @@ def search_graph(
"POST", f"{self._url_root}/api/v2/graph/search", json_data=params
)
return json.loads(response)

def search_agent_personas(
self,
name: str | None = None,
enabled: bool | None = None,
count: int = 50,
page: int = 0,
) -> list[YetiObject]:
"""Searches for agent personas in Yeti.

Unlike the other search methods, every argument is optional: listing all
personas is the common case.

Args:
name: The name of the persona to search for (substring match).
enabled: Restricts to enabled or disabled personas. None means both.
count: The number of results to return (default is 50).
page: The page of results to return (default is 0, which means the first page).

Returns:
The response from the API; a list of dicts representing personas.
"""
params: dict[str, Any] = {"name": name or "", "count": count, "page": page}
if enabled is not None:
params["enabled"] = enabled

response = self.do_request(
"POST",
f"{self._url_root}/api/v2/agentpersonas/search",
json_data=params,
)
return json.loads(response)["personas"]

def get_agent_persona(self, yeti_id: str) -> YetiObject:
"""Fetches a single agent persona by its Yeti ID.

Args:
yeti_id: The ID of the persona, as provided by Yeti.

Returns:
The response from the API; a dict representing the persona.
"""
response = self.do_request(
"GET", f"{self._url_root}/api/v2/agentpersonas/{yeti_id}"
)
return json.loads(response)

def new_agent_persona(self, persona: dict[str, Any]) -> YetiObject:
"""Creates a new agent persona in Yeti.

Args:
persona: The persona to create. Requires at least `name` and
`instruction`; Yeti rejects instructions shorter than 20 characters.

Returns:
The response from the API; a dict representing the persona.
"""
params = {"persona": persona}
response = self.do_request(
"POST",
f"{self._url_root}/api/v2/agentpersonas/",
json_data=params,
)
return json.loads(response)

def patch_agent_persona(self, yeti_id: str, persona: dict[str, Any]) -> YetiObject:
"""Updates an existing agent persona in Yeti.

Args:
yeti_id: The ID of the persona to update, as provided by Yeti.
persona: The full persona object to write.

Returns:
The response from the API; a dict representing the persona.
"""
params = {"persona": persona}
response = self.do_request(
"PATCH",
f"{self._url_root}/api/v2/agentpersonas/{yeti_id}",
json_data=params,
)
return json.loads(response)
Loading