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
11 changes: 11 additions & 0 deletions src/google/adk/models/google_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -453,6 +453,17 @@ async def connect(self, llm_request: LlmRequest) -> BaseLlmConnection:
' backend. Please use Vertex AI backend.'
)
llm_request.live_connect_config.tools = llm_request.config.tools
# Safety settings are configured via LlmAgent.generate_content_config, which
# only populates llm_request.config. Forward them so live runs honor the
# same safety configuration as non-live runs. An explicitly provided
# live_connect_config value takes precedence.
if (
llm_request.config.safety_settings is not None
and llm_request.live_connect_config.safety_settings is None
):
llm_request.live_connect_config.safety_settings = (
llm_request.config.safety_settings
)
logger.debug('Connecting to live with llm_request:%s', llm_request)
logger.debug('Live connect config: %s', llm_request.live_connect_config)
async with self._live_api_client.aio.live.connect(
Expand Down
144 changes: 144 additions & 0 deletions tests/unittests/models/test_google_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -852,6 +852,150 @@ async def __aexit__(self, *args):
)


@pytest.mark.asyncio
async def test_connect_forwards_safety_settings(gemini_llm, llm_request):
"""Live sessions receive safety_settings from generate_content_config."""
safety_settings = [
types.SafetySetting(
category=types.HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT,
threshold=types.HarmBlockThreshold.BLOCK_LOW_AND_ABOVE,
),
types.SafetySetting(
category=types.HarmCategory.HARM_CATEGORY_HARASSMENT,
threshold=types.HarmBlockThreshold.BLOCK_ONLY_HIGH,
),
]
llm_request.config.safety_settings = safety_settings
llm_request.live_connect_config = types.LiveConnectConfig()

mock_live_session = mock.AsyncMock()

with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:

class MockLiveConnect:

async def __aenter__(self):
return mock_live_session

async def __aexit__(self, *args):
pass

mock_live_client.aio.live.connect.return_value = MockLiveConnect()

async with gemini_llm.connect(llm_request) as connection:
mock_live_client.aio.live.connect.assert_called_once()
config_arg = mock_live_client.aio.live.connect.call_args.kwargs["config"]

assert config_arg.safety_settings == safety_settings
assert isinstance(connection, GeminiLlmConnection)


@pytest.mark.asyncio
async def test_connect_keeps_existing_live_safety_settings(
gemini_llm, llm_request
):
"""An explicit live_connect_config.safety_settings is not overwritten."""
live_safety_settings = [
types.SafetySetting(
category=types.HarmCategory.HARM_CATEGORY_HATE_SPEECH,
threshold=types.HarmBlockThreshold.BLOCK_NONE,
),
]
llm_request.config.safety_settings = [
types.SafetySetting(
category=types.HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT,
threshold=types.HarmBlockThreshold.BLOCK_LOW_AND_ABOVE,
),
]
llm_request.live_connect_config = types.LiveConnectConfig(
safety_settings=live_safety_settings
)

mock_live_session = mock.AsyncMock()

with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:

class MockLiveConnect:

async def __aenter__(self):
return mock_live_session

async def __aexit__(self, *args):
pass

mock_live_client.aio.live.connect.return_value = MockLiveConnect()

async with gemini_llm.connect(llm_request):
config_arg = mock_live_client.aio.live.connect.call_args.kwargs["config"]

assert config_arg.safety_settings == live_safety_settings


@pytest.mark.asyncio
async def test_connect_keeps_empty_live_safety_settings(
gemini_llm, llm_request
):
"""An explicit empty live_connect_config.safety_settings is not overwritten.

An empty list means "send no safety settings" and is distinct from None,
which means "not configured here".
"""
llm_request.config.safety_settings = [
types.SafetySetting(
category=types.HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT,
threshold=types.HarmBlockThreshold.BLOCK_LOW_AND_ABOVE,
),
]
llm_request.live_connect_config = types.LiveConnectConfig(safety_settings=[])

mock_live_session = mock.AsyncMock()

with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:

class MockLiveConnect:

async def __aenter__(self):
return mock_live_session

async def __aexit__(self, *args):
pass

mock_live_client.aio.live.connect.return_value = MockLiveConnect()

async with gemini_llm.connect(llm_request):
config_arg = mock_live_client.aio.live.connect.call_args.kwargs["config"]

assert config_arg.safety_settings is not None
assert len(config_arg.safety_settings) == 0


@pytest.mark.asyncio
async def test_connect_safety_settings_remain_none_when_unset(
gemini_llm, llm_request
):
"""No safety_settings anywhere leaves the live config untouched."""
llm_request.live_connect_config = types.LiveConnectConfig()

mock_live_session = mock.AsyncMock()

with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:

class MockLiveConnect:

async def __aenter__(self):
return mock_live_session

async def __aexit__(self, *args):
pass

mock_live_client.aio.live.connect.return_value = MockLiveConnect()

async with gemini_llm.connect(llm_request):
config_arg = mock_live_client.aio.live.connect.call_args.kwargs["config"]

assert config_arg.safety_settings is None


@pytest.mark.parametrize(
(
"api_backend, "
Expand Down
Loading