Make search_contacts interface more LLM-friendly (#272)
Break down `search_contacts` into `search_contacts_by_name` and `search_contacts_by_email`. The search_contacts' `query` argument was not clear enough for LLMs.
This commit is contained in:
parent
3f7226709f
commit
75da4bf8b0
5 changed files with 98 additions and 98 deletions
|
|
@ -16,3 +16,6 @@ except ValueError as e:
|
||||||
f"'{os.getenv('ARCADE_GMAIL_DEFAULT_REPLY_TO')}'. Expected one of "
|
f"'{os.getenv('ARCADE_GMAIL_DEFAULT_REPLY_TO')}'. Expected one of "
|
||||||
f"{list(GmailReplyToWhom.__members__.keys())}"
|
f"{list(GmailReplyToWhom.__members__.keys())}"
|
||||||
) from e
|
) from e
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_SEARCH_CONTACTS_LIMIT = 30
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,9 @@ from arcade.sdk.auth import Google
|
||||||
from google.oauth2.credentials import Credentials
|
from google.oauth2.credentials import Credentials
|
||||||
from googleapiclient.discovery import build
|
from googleapiclient.discovery import build
|
||||||
|
|
||||||
|
from arcade_google.tools.constants import DEFAULT_SEARCH_CONTACTS_LIMIT
|
||||||
|
from arcade_google.tools.utils import build_people_service, search_contacts
|
||||||
|
|
||||||
|
|
||||||
async def _warmup_cache(service) -> None: # type: ignore[no-untyped-def]
|
async def _warmup_cache(service) -> None: # type: ignore[no-untyped-def]
|
||||||
"""
|
"""
|
||||||
|
|
@ -18,47 +21,46 @@ async def _warmup_cache(service) -> None: # type: ignore[no-untyped-def]
|
||||||
|
|
||||||
|
|
||||||
@tool(requires_auth=Google(scopes=["https://www.googleapis.com/auth/contacts.readonly"]))
|
@tool(requires_auth=Google(scopes=["https://www.googleapis.com/auth/contacts.readonly"]))
|
||||||
async def search_contacts(
|
async def search_contacts_by_email(
|
||||||
context: ToolContext,
|
context: ToolContext,
|
||||||
query: Annotated[
|
email: Annotated[str, "The email address to search for"],
|
||||||
str,
|
|
||||||
"The search query for filtering contacts.",
|
|
||||||
],
|
|
||||||
limit: Annotated[
|
limit: Annotated[
|
||||||
Optional[int],
|
Optional[int],
|
||||||
"The maximum number of contacts to return (default 10, max 30)",
|
"The maximum number of contacts to return (30 is the max allowed by Google API)",
|
||||||
] = 10,
|
] = DEFAULT_SEARCH_CONTACTS_LIMIT,
|
||||||
) -> Annotated[dict, "A dictionary containing the list of matching contacts"]:
|
) -> Annotated[dict, "A dictionary containing the list of matching contacts"]:
|
||||||
"""
|
"""
|
||||||
Search the user's contacts in Google Contacts.
|
Search the user's contacts in Google Contacts by email address.
|
||||||
|
|
||||||
Up to 30 contacts with a name or email address containing the query will be returned.
|
|
||||||
If the query matches more than 30 contacts, only the first 30 will be returned.
|
|
||||||
"""
|
"""
|
||||||
# Build the People API service
|
service = build_people_service(
|
||||||
service = build(
|
context.authorization.token if context.authorization and context.authorization.token else ""
|
||||||
"people",
|
|
||||||
"v1",
|
|
||||||
credentials=Credentials(
|
|
||||||
context.authorization.token
|
|
||||||
if context.authorization and context.authorization.token
|
|
||||||
else ""
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Warm-up the cache before performing search.
|
# Warm-up the cache before performing search.
|
||||||
# TODO: Ideally we should warmup only if this user (or google domain?) hasn't warmed up recently
|
# TODO: Ideally we should warmup only if this user (or google domain?) hasn't warmed up recently
|
||||||
await _warmup_cache(service)
|
await _warmup_cache(service)
|
||||||
|
|
||||||
# Search primary contacts using searchContacts
|
return {"contacts": search_contacts(service, email, limit)}
|
||||||
primary_response = (
|
|
||||||
service.people()
|
|
||||||
.searchContacts(query=query, pageSize=limit, readMask="names,emailAddresses")
|
|
||||||
.execute()
|
|
||||||
)
|
|
||||||
primary_results = primary_response.get("results", [])
|
|
||||||
|
|
||||||
return {"contacts": primary_results}
|
|
||||||
|
@tool(requires_auth=Google(scopes=["https://www.googleapis.com/auth/contacts.readonly"]))
|
||||||
|
async def search_contacts_by_name(
|
||||||
|
context: ToolContext,
|
||||||
|
name: Annotated[str, "The full name to search for"],
|
||||||
|
limit: Annotated[
|
||||||
|
Optional[int],
|
||||||
|
"The maximum number of contacts to return (30 is the max allowed by Google API)",
|
||||||
|
] = DEFAULT_SEARCH_CONTACTS_LIMIT,
|
||||||
|
) -> Annotated[dict, "A dictionary containing the list of matching contacts"]:
|
||||||
|
"""
|
||||||
|
Search the user's contacts in Google Contacts by name.
|
||||||
|
"""
|
||||||
|
service = build_people_service(
|
||||||
|
context.authorization.token if context.authorization and context.authorization.token else ""
|
||||||
|
)
|
||||||
|
# Warm-up the cache before performing search.
|
||||||
|
# TODO: Ideally we should warmup only if this user (or google domain?) hasn't warmed up recently
|
||||||
|
await _warmup_cache(service)
|
||||||
|
return {"contacts": search_contacts(service, name, limit)}
|
||||||
|
|
||||||
|
|
||||||
@tool(requires_auth=Google(scopes=["https://www.googleapis.com/auth/contacts"]))
|
@tool(requires_auth=Google(scopes=["https://www.googleapis.com/auth/contacts"]))
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,7 @@ from datetime import datetime, timedelta
|
||||||
from email.message import EmailMessage
|
from email.message import EmailMessage
|
||||||
from email.mime.text import MIMEText
|
from email.mime.text import MIMEText
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, Optional, Union
|
from typing import Any, Optional, Union, cast
|
||||||
from zoneinfo import ZoneInfo
|
from zoneinfo import ZoneInfo
|
||||||
|
|
||||||
from arcade.sdk import ToolContext
|
from arcade.sdk import ToolContext
|
||||||
|
|
@ -13,6 +13,7 @@ from bs4 import BeautifulSoup
|
||||||
from google.oauth2.credentials import Credentials
|
from google.oauth2.credentials import Credentials
|
||||||
from googleapiclient.discovery import Resource, build
|
from googleapiclient.discovery import Resource, build
|
||||||
|
|
||||||
|
from arcade_google.tools.constants import DEFAULT_SEARCH_CONTACTS_LIMIT
|
||||||
from arcade_google.tools.exceptions import GmailToolError, GoogleServiceError
|
from arcade_google.tools.exceptions import GmailToolError, GoogleServiceError
|
||||||
from arcade_google.tools.models import Day, GmailAction, GmailReplyToWhom, TimeSlot
|
from arcade_google.tools.models import Day, GmailAction, GmailReplyToWhom, TimeSlot
|
||||||
|
|
||||||
|
|
@ -598,3 +599,39 @@ def build_docs_service(auth_token: Optional[str]) -> Resource: # type: ignore[n
|
||||||
"""
|
"""
|
||||||
auth_token = auth_token or ""
|
auth_token = auth_token or ""
|
||||||
return build("docs", "v1", credentials=Credentials(auth_token))
|
return build("docs", "v1", credentials=Credentials(auth_token))
|
||||||
|
|
||||||
|
|
||||||
|
# Contacts utils
|
||||||
|
def build_people_service(auth_token: Optional[str]) -> Resource: # type: ignore[no-any-unimported]
|
||||||
|
"""
|
||||||
|
Build a People service object.
|
||||||
|
"""
|
||||||
|
auth_token = auth_token or ""
|
||||||
|
return build("people", "v1", credentials=Credentials(auth_token))
|
||||||
|
|
||||||
|
|
||||||
|
def search_contacts(service: Any, query: str, limit: Optional[int]) -> list[dict[str, Any]]:
|
||||||
|
"""
|
||||||
|
Search the user's contacts in Google Contacts.
|
||||||
|
"""
|
||||||
|
response = (
|
||||||
|
service.people()
|
||||||
|
.searchContacts(
|
||||||
|
query=query,
|
||||||
|
pageSize=limit or DEFAULT_SEARCH_CONTACTS_LIMIT,
|
||||||
|
readMask=",".join([
|
||||||
|
"names",
|
||||||
|
"nicknames",
|
||||||
|
"emailAddresses",
|
||||||
|
"phoneNumbers",
|
||||||
|
"addresses",
|
||||||
|
"organizations",
|
||||||
|
"biographies",
|
||||||
|
"urls",
|
||||||
|
"userDefined",
|
||||||
|
]),
|
||||||
|
)
|
||||||
|
.execute()
|
||||||
|
)
|
||||||
|
|
||||||
|
return cast(list[dict[str, Any]], response.get("results", []))
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,11 @@ from arcade.sdk.eval import (
|
||||||
from arcade.sdk.eval.critic import BinaryCritic
|
from arcade.sdk.eval.critic import BinaryCritic
|
||||||
|
|
||||||
import arcade_google
|
import arcade_google
|
||||||
from arcade_google.tools.contacts import create_contact, search_contacts
|
from arcade_google.tools.contacts import (
|
||||||
|
create_contact,
|
||||||
|
search_contacts_by_email,
|
||||||
|
search_contacts_by_name,
|
||||||
|
)
|
||||||
|
|
||||||
# Evaluation rubric
|
# Evaluation rubric
|
||||||
rubric = EvalRubric(
|
rubric = EvalRubric(
|
||||||
|
|
@ -31,12 +35,27 @@ def contacts_eval_suite() -> EvalSuite:
|
||||||
)
|
)
|
||||||
|
|
||||||
suite.add_case(
|
suite.add_case(
|
||||||
name="Find a contact by name",
|
name="Search contacts by name",
|
||||||
user_message="Find my contact Bob",
|
user_message="Find my contact Bob",
|
||||||
expected_tool_calls=[
|
expected_tool_calls=[
|
||||||
ExpectedToolCall(
|
ExpectedToolCall(
|
||||||
func=search_contacts,
|
func=search_contacts_by_name,
|
||||||
args={"query": "Bob"},
|
args={
|
||||||
|
"name": "Bob",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
suite.add_case(
|
||||||
|
name="Search contacts by email",
|
||||||
|
user_message="Find my contact alice@example.com",
|
||||||
|
expected_tool_calls=[
|
||||||
|
ExpectedToolCall(
|
||||||
|
func=search_contacts_by_email,
|
||||||
|
args={
|
||||||
|
"email": "alice@example.com",
|
||||||
|
},
|
||||||
)
|
)
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
@ -46,9 +65,9 @@ def contacts_eval_suite() -> EvalSuite:
|
||||||
user_message="Find 5 contacts whose names include 'Alice'",
|
user_message="Find 5 contacts whose names include 'Alice'",
|
||||||
expected_tool_calls=[
|
expected_tool_calls=[
|
||||||
ExpectedToolCall(
|
ExpectedToolCall(
|
||||||
func=search_contacts,
|
func=search_contacts_by_name,
|
||||||
args={
|
args={
|
||||||
"query": "Alice",
|
"name": "Alice",
|
||||||
"limit": 5,
|
"limit": 5,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
import pytest
|
import pytest
|
||||||
from arcade.sdk import ToolContext
|
from arcade.sdk import ToolContext
|
||||||
|
|
||||||
from arcade_google.tools.contacts import create_contact, search_contacts
|
from arcade_google.tools.contacts import create_contact
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
|
|
@ -14,67 +14,6 @@ def mock_context():
|
||||||
return context
|
return context
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_search_contacts_success(mock_context):
|
|
||||||
search_response_data = {
|
|
||||||
"results": [
|
|
||||||
{
|
|
||||||
"resourceName": "people/1",
|
|
||||||
"names": [{"displayName": "John Doe"}],
|
|
||||||
"emailAddresses": [{"value": "john@example.com"}],
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"resourceName": "people/2",
|
|
||||||
"names": [{"displayName": "Jane Doe"}],
|
|
||||||
"emailAddresses": [{"value": "jane@example.com"}],
|
|
||||||
},
|
|
||||||
]
|
|
||||||
}
|
|
||||||
search_call = MagicMock()
|
|
||||||
search_call.execute.return_value = search_response_data
|
|
||||||
|
|
||||||
people_mock = MagicMock()
|
|
||||||
people_mock.searchContacts.return_value = search_call
|
|
||||||
|
|
||||||
service_mock = MagicMock()
|
|
||||||
service_mock.people.return_value = people_mock
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch("arcade_google.tools.contacts.build", return_value=service_mock) as mock_build,
|
|
||||||
patch(
|
|
||||||
"arcade_google.tools.contacts._warmup_cache", new=AsyncMock(return_value=None)
|
|
||||||
) as mock_warmup,
|
|
||||||
):
|
|
||||||
result = await search_contacts(mock_context, query="Doe", limit=2)
|
|
||||||
assert "contacts" in result
|
|
||||||
assert result["contacts"] == search_response_data["results"]
|
|
||||||
|
|
||||||
assert mock_warmup.call_count == 1
|
|
||||||
assert people_mock.searchContacts.call_count == 1
|
|
||||||
|
|
||||||
# Check that the People API service was built with the expected parameters.
|
|
||||||
mock_build.assert_called_once()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_search_contacts_error(mock_context):
|
|
||||||
error_call = MagicMock()
|
|
||||||
error_call.execute.side_effect = Exception("Search error")
|
|
||||||
|
|
||||||
people_mock = MagicMock()
|
|
||||||
people_mock.searchContacts.return_value = error_call
|
|
||||||
|
|
||||||
service_mock = MagicMock()
|
|
||||||
service_mock.people.return_value = people_mock
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch("arcade_google.tools.contacts.build", return_value=service_mock),
|
|
||||||
patch("arcade_google.tools.contacts._warmup_cache", new=AsyncMock(return_value=None)),
|
|
||||||
pytest.raises(Exception, match="Error in execution of SearchContacts"),
|
|
||||||
):
|
|
||||||
await search_contacts(mock_context, query="Doe")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_create_contact_success(mock_context):
|
async def test_create_contact_success(mock_context):
|
||||||
# Test create_contact with all parameters (given, family names and email)
|
# Test create_contact with all parameters (given, family names and email)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue