ADK changes

PiperOrigin-RevId: 829136628
This commit is contained in:
Google Team Member
2025-11-06 15:55:26 -08:00
committed by Copybara-Service
parent f3d6fcf444
commit f1f44675e4
16 changed files with 6481 additions and 50 deletions
+91 -29
View File
@@ -10,8 +10,7 @@
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# limitations under the Licens
import json
from unittest.mock import AsyncMock
@@ -91,6 +90,32 @@ LLM_REQUEST_WITH_FUNCTION_DECLARATION = LlmRequest(
),
)
FILE_URI_TEST_CASES = [
pytest.param("gs://bucket/document.pdf", "application/pdf", id="pdf"),
pytest.param("gs://bucket/data.json", "application/json", id="json"),
pytest.param("gs://bucket/data.txt", "text/plain", id="txt"),
]
FILE_BYTES_TEST_CASES = [
pytest.param(
b"test_pdf_data",
"application/pdf",
"data:application/pdf;base64,dGVzdF9wZGZfZGF0YQ==",
id="pdf",
),
pytest.param(
b'{"hello":"world"}',
"application/json",
"data:application/json;base64,eyJoZWxsbyI6IndvcmxkIn0=",
id="json",
),
pytest.param(
b"hello world",
"text/plain",
"data:text/plain;base64,aGVsbG8gd29ybGQ=",
id="txt",
),
]
STREAMING_MODEL_RESPONSE = [
ModelResponse(
@@ -1041,6 +1066,46 @@ def test_function_declaration_to_tool_param(
)
def test_function_declaration_to_tool_param_with_parameters_json_schema():
"""Ensure function declarations using parameters_json_schema are handled.
This verifies that when a FunctionDeclaration includes a raw
`parameters_json_schema` dict, it is used directly as the function
parameters in the resulting tool param.
"""
func_decl = types.FunctionDeclaration(
name="fn_with_json",
description="desc",
parameters_json_schema={
"type": "object",
"properties": {
"a": {"type": "string"},
"b": {"type": "array", "items": {"type": "string"}},
},
"required": ["a"],
},
)
expected = {
"type": "function",
"function": {
"name": "fn_with_json",
"description": "desc",
"parameters": {
"type": "object",
"properties": {
"a": {"type": "string"},
"b": {"type": "array", "items": {"type": "string"}},
},
"required": ["a"],
},
},
}
assert _function_declaration_to_tool_param(func_decl) == expected
@pytest.mark.asyncio
async def test_generate_content_async_with_system_instruction(
lite_llm_instance, mock_acompletion
@@ -1183,10 +1248,11 @@ def test_content_to_message_param_user_message():
assert message["content"] == "Test prompt"
def test_content_to_message_param_user_message_with_file_uri():
file_part = types.Part.from_uri(
file_uri="gs://bucket/document.pdf", mime_type="application/pdf"
)
@pytest.mark.parametrize("file_uri,mime_type", FILE_URI_TEST_CASES)
def test_content_to_message_param_user_message_with_file_uri(
file_uri, mime_type
):
file_part = types.Part.from_uri(file_uri=file_uri, mime_type=mime_type)
content = types.Content(
role="user",
parts=[
@@ -1201,14 +1267,15 @@ def test_content_to_message_param_user_message_with_file_uri():
assert message["content"][0]["type"] == "text"
assert message["content"][0]["text"] == "Summarize this file."
assert message["content"][1]["type"] == "file"
assert message["content"][1]["file"]["file_id"] == "gs://bucket/document.pdf"
assert message["content"][1]["file"]["file_id"] == file_uri
assert "format" not in message["content"][1]["file"]
def test_content_to_message_param_user_message_file_uri_only():
file_part = types.Part.from_uri(
file_uri="gs://bucket/only.pdf", mime_type="application/pdf"
)
@pytest.mark.parametrize("file_uri,mime_type", FILE_URI_TEST_CASES)
def test_content_to_message_param_user_message_file_uri_only(
file_uri, mime_type
):
file_part = types.Part.from_uri(file_uri=file_uri, mime_type=mime_type)
content = types.Content(
role="user",
parts=[
@@ -1220,7 +1287,7 @@ def test_content_to_message_param_user_message_file_uri_only():
assert message["role"] == "user"
assert isinstance(message["content"], list)
assert message["content"][0]["type"] == "file"
assert message["content"][0]["file"]["file_id"] == "gs://bucket/only.pdf"
assert message["content"][0]["file"]["file_id"] == file_uri
assert "format" not in message["content"][0]["file"]
@@ -1402,29 +1469,23 @@ def test_get_content_video():
assert "format" not in content[0]["video_url"]
def test_get_content_pdf():
parts = [
types.Part.from_bytes(data=b"test_pdf_data", mime_type="application/pdf")
]
@pytest.mark.parametrize(
"file_data,mime_type,expected_base64", FILE_BYTES_TEST_CASES
)
def test_get_content_file_bytes(file_data, mime_type, expected_base64):
parts = [types.Part.from_bytes(data=file_data, mime_type=mime_type)]
content = _get_content(parts)
assert content[0]["type"] == "file"
assert (
content[0]["file"]["file_data"]
== "data:application/pdf;base64,dGVzdF9wZGZfZGF0YQ=="
)
assert content[0]["file"]["file_data"] == expected_base64
assert "format" not in content[0]["file"]
def test_get_content_file_uri():
parts = [
types.Part.from_uri(
file_uri="gs://bucket/document.pdf",
mime_type="application/pdf",
)
]
@pytest.mark.parametrize("file_uri,mime_type", FILE_URI_TEST_CASES)
def test_get_content_file_uri(file_uri, mime_type):
parts = [types.Part.from_uri(file_uri=file_uri, mime_type=mime_type)]
content = _get_content(parts)
assert content[0]["type"] == "file"
assert content[0]["file"]["file_id"] == "gs://bucket/document.pdf"
assert content[0]["file"]["file_id"] == file_uri
assert "format" not in content[0]["file"]
@@ -1925,7 +1986,8 @@ async def test_generate_content_async_non_compliant_multiple_function_calls(
This test verifies that:
1. Multiple function calls with same indices (0) are handled correctly
2. Arguments and names are properly accumulated for each function call
3. The final response contains all function calls with correct incremented indices
3. The final response contains all function calls with correct incremented
indices
"""
mock_completion.return_value = NON_COMPLIANT_MULTIPLE_FUNCTION_CALLS_STREAM