Skip to content
Open
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
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import asyncio
import uuid
from abc import ABC, abstractmethod
from typing import Any, AsyncGenerator, Callable, Dict, List, Mapping, Sequence
from typing import Any, AsyncGenerator, Callable, Dict, List, Mapping, Sequence, Tuple

from autogen_core import (
AgentId,
Expand Down Expand Up @@ -67,7 +67,7 @@ def __init__(
self,
name: str,
description: str,
participants: List[ChatAgent | Team],
participants: List[ChatAgent | Team] | Tuple[ChatAgent | Team, ...],
group_chat_manager_name: str,
group_chat_manager_class: type[SequentialRoutedAgent],
termination_condition: TerminationCondition | None = None,
Expand All @@ -78,8 +78,19 @@ def __init__(
):
self._name = name
self._description = description
if participants is None:
raise TypeError("participants must be a list or tuple of ChatAgent or Team instances, got None")
if not isinstance(participants, (list, tuple)):
raise TypeError(
f"participants must be a list or tuple of ChatAgent or Team instances, got {type(participants).__name__}"
)
if len(participants) == 0:
raise ValueError("At least one participant is required.")
for i, participant in enumerate(participants):
if not isinstance(participant, (ChatAgent, Team)):
raise TypeError(
f"participants[{i}] must be a ChatAgent or Team instance, got {type(participant).__name__}"
)
if len(participants) != len(set(participant.name for participant in participants)):
raise ValueError("The participant names must be unique.")
self._participants = participants
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import asyncio
from typing import Any, Callable, List, Mapping, Sequence
from typing import Any, Callable, List, Mapping, Sequence, Tuple

from autogen_core import AgentRuntime, Component, ComponentModel
from pydantic import BaseModel
Expand Down Expand Up @@ -241,7 +241,7 @@ async def main() -> None:

def __init__(
self,
participants: List[ChatAgent | Team],
participants: List[ChatAgent | Team] | Tuple[ChatAgent | Team, ...],
*,
name: str | None = None,
description: str | None = None,
Expand Down
32 changes: 32 additions & 0 deletions python/packages/autogen-agentchat/tests/test_group_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -1944,3 +1944,35 @@ async def test_selector_group_chat_streaming(runtime: AgentRuntime | None) -> No

# Content-based verification instead of index-based
# Note: The streaming test verifies the streaming behavior, not the final result content


@pytest.mark.asyncio
async def test_round_robin_group_chat_validates_participants_none() -> None:
"""Test that participants=None raises a clear TypeError."""
with pytest.raises(TypeError, match="participants must be a list or tuple"):
RoundRobinGroupChat(participants=None) # type: ignore


@pytest.mark.asyncio
async def test_round_robin_group_chat_validates_participants_not_sequence() -> None:
"""Test that non-sequence participants raises a clear TypeError."""
with pytest.raises(TypeError, match="participants must be a list or tuple"):
RoundRobinGroupChat(participants="not a list") # type: ignore


@pytest.mark.asyncio
async def test_round_robin_group_chat_validates_participants_invalid_type() -> None:
"""Test that participants containing non-agent objects raises a clear TypeError."""
from autogen_agentchat.agents import UserProxyAgent

with pytest.raises(TypeError, match=r"participants\[1\] must be a ChatAgent or Team"):
RoundRobinGroupChat(participants=[UserProxyAgent(name="valid"), "invalid"]) # type: ignore


@pytest.mark.asyncio
async def test_round_robin_group_chat_accepts_tuple_participants() -> None:
"""Test that participants passed as a tuple is accepted."""
from autogen_agentchat.agents import UserProxyAgent

team = RoundRobinGroupChat(participants=(UserProxyAgent(name="agent1"), UserProxyAgent(name="agent2")))
assert len(team._participants) == 2