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
8 changes: 5 additions & 3 deletions src/agents/extensions/memory/advanced_sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -1012,8 +1012,10 @@ async def create_branch_from_turn(

Raises:
ValueError: If turn doesn't exist, doesn't contain a user message, or
`branch_name` has already been used in this session
`branch_name` is blank or has already been used in this session.
"""
if branch_name is not None and not branch_name.strip():
raise ValueError("Branch name cannot be empty")

async def _create_and_switch() -> tuple[str, Any, str]:
# Copying the branch is the first durable side effect. Keep the
Expand Down Expand Up @@ -1066,8 +1068,8 @@ async def create_branch_from_content(
The branch_id of the newly created branch.

Raises:
ValueError: If no matching turns are found or `branch_name` has already been used
in this session.
ValueError: If no matching turns are found or `branch_name` is blank or has
already been used in this session.
"""
matching_turns = await self.find_turns_by_content(search_term)
if not matching_turns:
Expand Down
27 changes: 27 additions & 0 deletions tests/extensions/memory/test_advanced_sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -1322,6 +1322,33 @@ def _pop_item_in_process(
session.close()


@pytest.mark.parametrize("branch_name", ["", " ", "\t"])
@pytest.mark.parametrize("from_content", [False, True])
async def test_create_branch_rejects_blank_name(branch_name: str, from_content: bool):
session = AdvancedSQLiteSession(session_id="blank_branch_name", create_tables=True)
items: list[TResponseInputItem] = [
{"role": "user", "content": "First question"},
{"role": "assistant", "content": "First answer"},
{"role": "user", "content": "Second question"},
]

try:
await session.add_items(items)
branches_before = await session.list_branches()

with pytest.raises(ValueError, match="Branch name cannot be empty"):
if from_content:
await session.create_branch_from_content("Second question", branch_name)
else:
await session.create_branch_from_turn(2, branch_name)

assert await session.list_branches() == branches_before
assert session._current_branch_id == "main"
assert await session.get_items() == items
finally:
session.close()


@pytest.mark.parametrize("branch_id", ["main", "existing_branch"])
async def test_create_branch_rejects_populated_branch_id(branch_id: str):
"""Creating a branch must not append history to a populated branch."""
Expand Down
Loading