diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/_utils.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/_utils.py new file mode 100644 index 00000000..e7bf85ea --- /dev/null +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/_utils.py @@ -0,0 +1,27 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +from logging import Logger + +import aiohttp + + +def _handle_request_error( + logger: Logger, response: aiohttp.ClientResponse, resource: str = "resource" +) -> None: + + if response.status == 400: + logger.error("Bad request for '%s': %s", resource, response.status) + else: + logger.error("Error accessing '%s': %s", resource, response.status) + + if not response.ok: + response.raise_for_status() + + raise aiohttp.ClientResponseError( + response.request_info, + response.history, + status=response.status, + message=f"Error accessing resource '{resource}'", + headers=response.headers, + ) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/agent_sign_in.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/agent_sign_in.py index 31b40754..2df8feca 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/agent_sign_in.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/agent_sign_in.py @@ -7,6 +7,7 @@ from microsoft_agents.activity import SignInResource from ..telemetry import user_token_client_spans as spans from ..agent_sign_in_base import AgentSignInBase +from .._utils import _handle_request_error from ._base_client import _BaseClient logger = logging.getLogger(__name__) @@ -60,9 +61,10 @@ async def get_sign_in_url( async with self._wrapped_client().get( "api/agentsignin/getSignInUrl", params=params ) as response: - if response.status >= 300: - logger.error("Error getting sign-in URL: %s", response.status) - response.raise_for_status() + if response.status != 200: + _handle_request_error( + logger, response, resource="api/agentsignin/getSignInUrl" + ) return await response.text() @@ -100,9 +102,10 @@ async def get_sign_in_resource( "api/botsignin/getSignInResource", params=params ) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error("Error getting sign-in resource: %s", response.status) - response.raise_for_status() + if response.status != 200: + _handle_request_error( + logger, response, resource="api/botsignin/getSignInResource" + ) data = await response.json() return SignInResource.model_validate(data) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/connector_client.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/connector_client.py index 1847e660..e969ef89 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/connector_client.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/connector_client.py @@ -11,12 +11,14 @@ from microsoft_agents.activity import ( Activity, + AttachmentData, ChannelAccount, Channels, ConversationParameters, ConversationResourceResponse, ResourceResponse, RoleTypes, + Transcript, ConversationsResult, PagedMembersResult, ) @@ -26,6 +28,7 @@ from ..conversations_base import ConversationsBase from ..get_product_info import get_product_info from ..telemetry import connector_spans as spans +from .._utils import _handle_request_error from ._base_client import _BaseClient logger = logging.getLogger(__name__) @@ -40,16 +43,6 @@ def __init__(self, **kwargs): self.views = kwargs.get("views") -class AttachmentData: - """Data for an attachment.""" - - def __init__(self, **kwargs): - self.name = kwargs.get("name") - self.original_base64 = kwargs.get("originalBase64") - self.type = kwargs.get("type") - self.thumbnail_base64 = kwargs.get("thumbnailBase64") - - def normalize_outgoing_activity(data: Any) -> Any: """ Normalizes an outgoing activity object for wire transmission. @@ -96,13 +89,10 @@ async def get_attachment_info(self, attachment_id: str) -> AttachmentInfo: async with self._wrapped_client().get(url) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error( - "Error getting attachment info: %s", - response.status, - stack_info=True, + if response.status != 200: + _handle_request_error( + logger, response, resource=f"v3/attachments/" ) - response.raise_for_status() data = await response.json() return AttachmentInfo(**data) @@ -139,11 +129,19 @@ async def get_attachment(self, attachment_id: str, view_id: str) -> BytesIO: async with self._wrapped_client().get(url) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error( - "Error getting attachment: %s", response.status, stack_info=True + if response.status in (301, 302): + logger.warning( + "Redirect when getting attachment: %s", + response.status, + stack_info=True, + ) + return BytesIO() + elif response.status >= 300: + _handle_request_error( + logger, + response, + resource=f"v3/attachments//views/", ) - response.raise_for_status() data = await response.read() return BytesIO(data) @@ -218,13 +216,8 @@ async def get_conversations( ) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error( - "Error getting conversations: %s", - response.status, - stack_info=True, - ) - response.raise_for_status() + if response.status != 200: + _handle_request_error(logger, response, resource="v3/conversations") data = await response.json() return ConversationsResult.model_validate(data) @@ -246,13 +239,9 @@ async def create_conversation( json=body.model_dump(by_alias=True, exclude_unset=True, mode="json"), ) as response: span.share(http_method="POST", status_code=response.status) - if response.status >= 300: - logger.error( - "Error creating conversation: %s", - response.status, - stack_info=True, - ) - response.raise_for_status() + + if response.status not in (200, 201, 202): + _handle_request_error(logger, response, resource="v3/conversations") data = await response.json() return ConversationResourceResponse.model_validate(data) @@ -297,13 +286,12 @@ async def reply_to_activity( response_text = await response.text("utf-8") - if response.status >= 300: - logger.error( - "Error replying to activity: %s", - response_text or response.status, - stack_info=True, + if response.status not in (200, 201, 202): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/activities/{activity_id}", ) - response.raise_for_status() if not response_text: resource_response = ResourceResponse() @@ -354,13 +342,12 @@ async def send_to_conversation( ) as response: span.share(http_method="POST", status_code=response.status) - if response.status >= 300: - logger.error( - "Error sending to conversation: %s", - response.status, - stack_info=True, + if response.status not in (200, 201, 202): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/activities", ) - response.raise_for_status() response_text = await response.text("utf-8") if not response_text: @@ -401,11 +388,12 @@ async def update_activity( url, json=body.model_dump(by_alias=True, exclude_unset=True), ) as response: - if response.status >= 300: - logger.error( - "Error updating activity: %s", response.status, stack_info=True + if response.status not in (200, 201, 202): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/activities/{activity_id}", ) - response.raise_for_status() data = await response.json() return ResourceResponse.model_validate(data) @@ -438,11 +426,12 @@ async def delete_activity(self, conversation_id: str, activity_id: str) -> None: async with self._wrapped_client().delete(url) as response: span.share(http_method="DELETE", status_code=response.status) - if response.status >= 300: - logger.error( - "Error deleting activity: %s", response.status, stack_info=True + if response.status not in (200, 202): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/activities/{activity_id}", ) - response.raise_for_status() async def upload_attachment( self, conversation_id: str, body: AttachmentData @@ -466,14 +455,6 @@ async def upload_attachment( conversation_id = self._normalize_conversation_id(conversation_id) url = f"v3/conversations/{conversation_id}/attachments" - # Convert the AttachmentData to a dictionary - attachment_dict = { - "name": body.name, - "originalBase64": body.original_base64, - "type": body.type, - "thumbnailBase64": body.thumbnail_base64, - } - logger.info( "Uploading attachment to conversation: %s, Attachment name: %s", conversation_id, @@ -481,17 +462,17 @@ async def upload_attachment( ) async with self._wrapped_client().post( - url, json=attachment_dict + url, + json=body.model_dump(by_alias=True, exclude_unset=True, mode="json"), ) as response: span.share(http_method="POST", status_code=response.status) - if response.status >= 300: - logger.error( - "Error uploading attachment: %s", - response.status, - stack_info=True, + if response.status not in (200, 201, 202): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/attachments", ) - response.raise_for_status() data = await response.json() return ResourceResponse.model_validate(data) @@ -524,13 +505,12 @@ async def get_conversation_members( async with self._wrapped_client().get(url) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error( - "Error getting conversation members: %s", - response.status, - stack_info=True, + if response.status != 200: + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/members", ) - response.raise_for_status() data = await response.json() return [ChannelAccount.model_validate(member) for member in data] @@ -566,13 +546,12 @@ async def get_conversation_member( async with self._wrapped_client().get(url) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error( - "Error getting conversation member: %s", - response.status, - stack_info=True, + if response.status != 200: + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/members/{member_id}", ) - response.raise_for_status() data = await response.json() return ChannelAccount.model_validate(data) @@ -603,13 +582,12 @@ async def delete_conversation_member( ) async with self._wrapped_client().delete(url) as response: - if response.status >= 300: - logger.error( - "Error deleting conversation member: %s", - response.status, - stack_info=True, + if response.status not in (200, 204): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/members/{member_id}", ) - response.raise_for_status() async def get_activity_members( self, conversation_id: str, activity_id: str @@ -638,13 +616,12 @@ async def get_activity_members( ) async with self._wrapped_client().get(url) as response: - if response.status >= 300: - logger.error( - "Error getting activity members: %s", - response.status, - stack_info=True, + if response.status != 200: + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/activities/{activity_id}/members", ) - response.raise_for_status() data = await response.json() return [ChannelAccount.model_validate(member) for member in data] @@ -687,19 +664,18 @@ async def get_conversation_paged_members( ) async with self._wrapped_client().get(url, params=params) as response: - if response.status >= 300: - logger.error( - "Error getting conversation paged members: %s", - response.status, - stack_info=True, + if response.status != 200: + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/pagedmembers", ) - response.raise_for_status() data = await response.json() return PagedMembersResult.model_validate(data) async def send_conversation_history( - self, conversation_id: str, body: Any + self, conversation_id: str, body: Transcript ) -> ResourceResponse: """ Sends conversation history to a conversation. @@ -719,14 +695,18 @@ async def send_conversation_history( url = f"v3/conversations/{conversation_id}/activities/history" logger.info("Sending conversation history to conversation: %s", conversation_id) - async with self._wrapped_client().post(url, json=body) as response: - if response.status >= 300: - logger.error( - "Error sending conversation history: %s", - response.status, - stack_info=True, + async with self._wrapped_client().post( + url, + json=body.model_dump( + by_alias=True, exclude_unset=True, exclude_none=True, mode="json" + ), + ) as response: + if response.status not in (200, 201, 202): + _handle_request_error( + logger, + response, + resource=f"v3/conversations/{conversation_id}/activities/history", ) - response.raise_for_status() data = await response.json() return ResourceResponse.model_validate(data) diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/user_token.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/user_token.py index 74c01afc..20e7da26 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/user_token.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/client/user_token.py @@ -13,6 +13,7 @@ ) from ..telemetry import user_token_client_spans as spans from ..user_token_base import UserTokenBase +from .._utils import _handle_request_error from ._base_client import _BaseClient logger = logging.getLogger(__name__) @@ -67,9 +68,15 @@ async def get_token( ) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error("Error getting token: %s", response.status) - response.raise_for_status() + if response.status == 404: + logger.warning( + "404: Could be issue with magic code or user not found. Returning empty token response." + ) + return TokenResponse() + elif response.status != 200: + _handle_request_error( + logger, response, resource="api/usertoken/GetToken" + ) data = await response.json() return TokenResponse.model_validate(data) @@ -108,10 +115,11 @@ async def _get_token_or_sign_in_resource( span.share(http_method="GET", status_code=response.status) if response.status != 200: - logger.error( - "Error getting token or sign-in resource: %s", response.status + _handle_request_error( + logger, + response, + resource="/api/usertoken/GetTokenOrSignInResource", ) - response.raise_for_status() data = await response.json() return TokenOrSignInResourceResponse.model_validate(data) @@ -141,9 +149,10 @@ async def get_aad_tokens( ) as response: span.share(http_method="POST", status_code=response.status) - if response.status >= 300: - logger.error("Error getting AAD tokens: %s", response.status) - response.raise_for_status() + if response.status != 200: + _handle_request_error( + logger, response, resource="api/usertoken/GetAadTokens" + ) data = await response.json() return {k: TokenResponse.model_validate(v) for k, v in data.items()} @@ -174,9 +183,10 @@ async def sign_out( ) as response: span.share(http_method="DELETE", status_code=response.status) - if response.status >= 300: - logger.error("Error signing out: %s", response.status) - response.raise_for_status() + if response.status not in (200, 204): + _handle_request_error( + logger, response, resource="api/usertoken/SignOut" + ) async def get_token_status( self, @@ -204,9 +214,10 @@ async def get_token_status( ) as response: span.share(http_method="GET", status_code=response.status) - if response.status >= 300: - logger.error("Error getting token status: %s", response.status) - response.raise_for_status() + if response.status != 200: + _handle_request_error( + logger, response, resource="api/usertoken/GetTokenStatus" + ) data = await response.json() return [TokenStatus.model_validate(status) for status in data] diff --git a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/mcs/mcs_connector_client.py b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/mcs/mcs_connector_client.py index fc8bd000..95ddd4d5 100644 --- a/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/mcs/mcs_connector_client.py +++ b/libraries/microsoft-agents-hosting-core/microsoft_agents/hosting/core/connector/mcs/mcs_connector_client.py @@ -8,10 +8,12 @@ from aiohttp import ClientSession from microsoft_agents.activity import Activity, ResourceResponse + from ..connector_client_base import ConnectorClientBase from ..attachments_base import AttachmentsBase from ..conversations_base import ConversationsBase from ..client._base_client import _BaseClient +from .._utils import _handle_request_error logger = logging.getLogger(__name__) @@ -54,13 +56,9 @@ async def send_to_conversation( json=activity.model_dump(by_alias=True, exclude_unset=True, mode="json"), headers={"Accept": "application/json", "Content-Type": "application/json"}, ) as response: - if response.status >= 300: - logger.error( - "MCS Connector: Error sending activity: %s", - response.status, - stack_info=True, - ) - response.raise_for_status() + + if response.status not in (200, 201, 202): + _handle_request_error(logger, response, resource=self._endpoint) data = await response.json() return ResourceResponse.model_validate(data) diff --git a/tests/hosting_core/connector/test_connector_client.py b/tests/hosting_core/connector/test_connector_client.py index d70e9cd6..8659c320 100644 --- a/tests/hosting_core/connector/test_connector_client.py +++ b/tests/hosting_core/connector/test_connector_client.py @@ -3,14 +3,20 @@ """Tests for ConversationsOperations using aiohttp TestServer.""" -import json - import pytest -from aiohttp import web, ClientSession +from aiohttp import ClientResponseError, web, ClientSession from aiohttp.test_utils import TestServer -from microsoft_agents.activity import Activity, Channels, ResourceResponse, RoleTypes -from microsoft_agents.activity.channel_account import ChannelAccount +from microsoft_agents.activity import ( + Activity, + AttachmentData, + ChannelAccount, + Channels, + ConversationParameters, + ResourceResponse, + RoleTypes, + Transcript, +) from microsoft_agents.hosting.core.connector.client.connector_client import ( ConnectorClient, ConversationsOperations, @@ -128,6 +134,273 @@ async def handler(request): await server.close() +class TestConnectorClientContract: + """Tests the externally observable Connector REST contract.""" + + @pytest.mark.asyncio + async def test_conversation_lifecycle_uses_expected_requests_and_models(self): + requests = [] + + async def handler(request): + body = await request.json() if request.can_read_body else None + requests.append( + { + "method": request.method, + "path": request.path, + "query": dict(request.query), + "body": body, + } + ) + if request.method == "GET": + return web.json_response( + {"continuationToken": "next-page", "conversations": []} + ) + if request.path == "/v3/conversations": + return web.json_response({"id": "new-conversation"}, status=201) + if request.method == "DELETE": + return web.Response(status=202) + return web.json_response({"id": "activity-result"}, status=202) + + app = _create_app([web.route("*", "/{tail:.*}", handler)]) + server = TestServer(app) + await server.start_server() + client = ConnectorClient(str(server.make_url("/")), token="") + try: + conversations = await client.conversations.get_conversations("page-1") + created = await client.conversations.create_conversation( + ConversationParameters( + is_group=True, + bot=ChannelAccount(id="agent-1"), + members=[ChannelAccount(id="member-1")], + topic_name="Planning", + ) + ) + updated = await client.conversations.update_activity( + "conversation-1", + "activity-1", + Activity(type="message", text="updated"), + ) + await client.conversations.delete_activity("conversation-1", "activity-1") + history = await client.conversations.send_conversation_history( + "conversation-1", + Transcript( + activities=[Activity(type="message", text="Earlier message")] + ), + ) + finally: + await client.close() + await server.close() + + assert conversations.continuation_token == "next-page" + assert created.id == "new-conversation" + assert updated.id == "activity-result" + assert history.id == "activity-result" + assert requests == [ + { + "method": "GET", + "path": "/v3/conversations", + "query": {"continuationToken": "page-1"}, + "body": None, + }, + { + "method": "POST", + "path": "/v3/conversations", + "query": {}, + "body": { + "isGroup": True, + "bot": {"id": "agent-1"}, + "members": [{"id": "member-1"}], + "topicName": "Planning", + }, + }, + { + "method": "PUT", + "path": "/v3/conversations/conversation-1/activities/activity-1", + "query": {}, + "body": {"type": "message", "text": "updated"}, + }, + { + "method": "DELETE", + "path": "/v3/conversations/conversation-1/activities/activity-1", + "query": {}, + "body": None, + }, + { + "method": "POST", + "path": "/v3/conversations/conversation-1/activities/history", + "query": {}, + "body": { + "activities": [{"type": "message", "text": "Earlier message"}] + }, + }, + ] + + @pytest.mark.asyncio + async def test_members_paging_and_upload_use_expected_requests(self): + requests = [] + + async def handler(request): + body = await request.json() if request.can_read_body else None + requests.append((request.method, request.path, dict(request.query), body)) + if request.path.endswith("/pagedmembers"): + return web.json_response( + { + "continuationToken": "next", + "members": [{"id": "member-2", "name": "Ada"}], + } + ) + if request.path.endswith("/members/member-1"): + if request.method == "DELETE": + return web.Response(status=204) + return web.json_response({"id": "member-1", "name": "Lin"}) + if request.path.endswith("/members"): + return web.json_response([{"id": "member-1", "name": "Lin"}]) + if request.path.endswith("/attachments"): + return web.json_response({"id": "attachment-1"}) + raise AssertionError(f"Unexpected request: {request.method} {request.path}") + + app = _create_app([web.route("*", "/{tail:.*}", handler)]) + server = TestServer(app) + await server.start_server() + client = ConnectorClient(str(server.make_url("/")), token="") + try: + members = await client.conversations.get_conversation_members( + "conversation-1" + ) + member = await client.conversations.get_conversation_member( + "conversation-1", "member-1" + ) + page = await client.conversations.get_conversation_paged_members( + "conversation-1", page_size=25, continuation_token="page-1" + ) + await client.conversations.delete_conversation_member( + "conversation-1", "member-1" + ) + attachment = AttachmentData( + name="notes.txt", + type="text/plain", + original_base64=b"dGVzdA==", + ) + uploaded = await client.conversations.upload_attachment( + "conversation-1", + attachment, + ) + finally: + await client.close() + await server.close() + + assert members[0].id == "member-1" + assert member.name == "Lin" + assert page.members[0].id == "member-2" + assert uploaded.id == "attachment-1" + assert requests == [ + ( + "GET", + "/v3/conversations/conversation-1/members", + {}, + None, + ), + ( + "GET", + "/v3/conversations/conversation-1/members/member-1", + {}, + None, + ), + ( + "GET", + "/v3/conversations/conversation-1/pagedmembers", + {"pageSize": "25", "continuationToken": "page-1"}, + None, + ), + ( + "DELETE", + "/v3/conversations/conversation-1/members/member-1", + {}, + None, + ), + ( + "POST", + "/v3/conversations/conversation-1/attachments", + {}, + { + "name": "notes.txt", + "originalBase64": "dGVzdA==", + "type": "text/plain", + }, + ), + ] + + @pytest.mark.asyncio + async def test_attachment_operations_return_metadata_and_binary_content(self): + async def info_handler(request): + return web.json_response( + { + "name": "report.pdf", + "type": "application/pdf", + "views": [{"viewId": "original"}], + } + ) + + async def content_handler(request): + return web.Response(body=b"attachment bytes") + + app = _create_app( + [ + web.get("/v3/attachments/{attachment_id}", info_handler), + web.get( + "/v3/attachments/{attachment_id}/views/{view_id}", + content_handler, + ), + ] + ) + server = TestServer(app) + await server.start_server() + client = ConnectorClient(str(server.make_url("/")), token="") + try: + info = await client.attachments.get_attachment_info("attachment-1") + content = await client.attachments.get_attachment( + "attachment-1", "original" + ) + finally: + await client.close() + await server.close() + + assert info.name == "report.pdf" + assert info.type == "application/pdf" + assert info.views == [{"viewId": "original"}] + assert content.read() == b"attachment bytes" + + @pytest.mark.asyncio + @pytest.mark.parametrize("status", [302, 400, 500]) + async def test_unexpected_response_status_raises_client_response_error( + self, status + ): + async def handler(request): + return web.Response(status=status) + + app = _create_app( + [web.post("/v3/conversations/{conversation_id}/activities", handler)] + ) + server = TestServer(app) + await server.start_server() + client = ConnectorClient(str(server.make_url("/")), token="") + try: + with pytest.raises(ClientResponseError) as exc_info: + await client.conversations.send_to_conversation( + "conversation-1", Activity(type="message", text="hello") + ) + finally: + await client.close() + await server.close() + + assert exc_info.value.status == status + if status == 302: + assert ( + exc_info.value.message == "Error accessing resource " + "'v3/conversations/conversation-1/activities'" + ) + + class TestReplyToActivity: """Tests for ConversationsOperations.reply_to_activity.""" diff --git a/tests/hosting_core/connector/test_user_token_client.py b/tests/hosting_core/connector/test_user_token_client.py index c5e726f8..ff08f7fa 100644 --- a/tests/hosting_core/connector/test_user_token_client.py +++ b/tests/hosting_core/connector/test_user_token_client.py @@ -3,10 +3,20 @@ """Tests for UserToken Bot Framework operations.""" +import base64 +import json + import pytest -from aiohttp import ClientSession, web +from aiohttp import ClientResponseError, ClientSession, web from aiohttp.test_utils import TestServer +from microsoft_agents.activity import ( + Activity, + ChannelAccount, + ChannelId, + ConversationAccount, + TokenExchangeRequest, +) from microsoft_agents.hosting.core.connector.client.user_token_client import ( UserToken, UserTokenClient, @@ -142,3 +152,258 @@ async def handler(request): finally: await client.close() await server.close() + + +class TestUserTokenClientContract: + """Tests the externally observable OAuth token service contract.""" + + @pytest.mark.asyncio + async def test_token_operations_use_expected_requests_and_typed_responses(self): + requests = [] + + async def handler(request): + body = await request.json() if request.can_read_body else None + requests.append((request.method, request.path, dict(request.query), body)) + if request.path.endswith("/GetToken"): + return web.json_response( + {"token": "user-token", "connectionName": "connection"} + ) + if request.path.endswith("/SignOut"): + return web.Response(status=204) + if request.path.endswith("/GetTokenStatus"): + return web.json_response( + [{"connectionName": "connection", "hasToken": True}] + ) + if request.path.endswith("/GetAadTokens"): + return web.json_response( + {"https://graph.microsoft.com": {"token": "graph-token"}} + ) + if request.path.endswith("/exchange"): + return web.json_response({"token": "exchanged-token"}) + raise AssertionError(f"Unexpected token operation: {request.path}") + + app = web.Application() + app.router.add_route("*", "/{tail:.*}", handler) + server = TestServer(app) + await server.start_server() + client = UserTokenClient(str(server.make_url("/")), token="", app_id="app-id") + try: + token = await client.get_user_token( + "user-1", "connection", "msteams:copilot", "magic-code" + ) + await client.sign_out_user("user-1", "connection", "msteams:copilot") + statuses = await client.get_token_status( + "user-1", "msteams:copilot", include="configured" + ) + aad_tokens = await client.get_aad_tokens( + "user-1", + "connection", + ["https://graph.microsoft.com"], + "msteams:copilot", + ) + exchanged = await client.exchange_token( + "user-1", + "connection", + "msteams:copilot", + TokenExchangeRequest(uri="api://resource", token="subject-token"), + ) + finally: + await client.close() + await server.close() + + assert token.token == "user-token" + assert statuses[0].connection_name == "connection" + assert statuses[0].has_token is True + assert aad_tokens["https://graph.microsoft.com"].token == "graph-token" + assert exchanged.token == "exchanged-token" + assert requests == [ + ( + "GET", + "/api/usertoken/GetToken", + { + "userId": "user-1", + "connectionName": "connection", + "channelId": "msteams", + "code": "magic-code", + }, + None, + ), + ( + "DELETE", + "/api/usertoken/SignOut", + { + "userId": "user-1", + "connectionName": "connection", + "channelId": "msteams", + }, + None, + ), + ( + "GET", + "/api/usertoken/GetTokenStatus", + { + "userId": "user-1", + "channelId": "msteams", + "include": "configured", + }, + None, + ), + ( + "POST", + "/api/usertoken/GetAadTokens", + { + "userId": "user-1", + "connectionName": "connection", + "channelId": "msteams", + }, + {"resourceUrls": ["https://graph.microsoft.com"]}, + ), + ( + "POST", + "/api/usertoken/exchange", + { + "userId": "user-1", + "connectionName": "connection", + "channelId": "msteams", + }, + {"uri": "api://resource", "token": "subject-token"}, + ), + ] + + @pytest.mark.asyncio + async def test_sign_in_requests_carry_activity_context_in_encoded_state(self): + requests = [] + + async def handler(request): + requests.append((request.path, dict(request.query))) + if request.path.endswith("/getSignInResource"): + return web.json_response({"signInLink": "https://sign-in.example"}) + return web.json_response({"tokenResponse": {"token": "existing-token"}}) + + app = web.Application() + app.router.add_route("*", "/{tail:.*}", handler) + server = TestServer(app) + await server.start_server() + client = UserTokenClient(str(server.make_url("/")), token="", app_id="app-id") + activity = Activity( + type="message", + id="activity-1", + channel_id=ChannelId("msteams:copilot"), + service_url="https://service.example", + from_property=ChannelAccount(id="user-1"), + recipient=ChannelAccount(id="agent-1"), + conversation=ConversationAccount(id="conversation-1"), + ) + try: + sign_in = await client.get_sign_in_resource( + "connection", activity, final_redirect="https://app.example/done" + ) + token_or_sign_in = await client.get_token_or_sign_in_resource( + "connection", + activity, + code="magic-code", + final_redirect="https://app.example/done", + fwd_url="https://app.example/forward", + ) + finally: + await client.close() + await server.close() + + assert sign_in.sign_in_link == "https://sign-in.example" + assert token_or_sign_in.token_response.token == "existing-token" + + sign_in_path, sign_in_query = requests[0] + token_path, token_query = requests[1] + sign_in_state = json.loads( + base64.b64decode(sign_in_query.pop("state")).decode() + ) + token_state = json.loads(base64.b64decode(token_query.pop("state")).decode()) + expected_state = { + "connectionName": "connection", + "conversation": { + "activityId": "activity-1", + "user": {"id": "user-1"}, + "bot": {"id": "agent-1"}, + "conversation": {"id": "conversation-1"}, + "channelId": "msteams", + "serviceUrl": "https://service.example", + }, + "msAppId": "app-id", + } + assert sign_in_path == "/api/botsignin/getSignInResource" + assert sign_in_query == {"finalRedirect": "https://app.example/done"} + assert sign_in_state == expected_state + assert token_path == "/api/usertoken/GetTokenOrSignInResource" + assert token_query == { + "userId": "user-1", + "connectionName": "connection", + "channelId": "msteams", + "code": "magic-code", + "finalRedirect": "https://app.example/done", + "fwdUrl": "https://app.example/forward", + } + assert token_state == expected_state + + @pytest.mark.asyncio + async def test_missing_context_is_rejected_before_a_sign_in_request(self): + client_without_app_id = UserTokenClient( + "https://token.example", token="", app_id=None + ) + activity = Activity( + type="message", + from_property=ChannelAccount(id="user-1"), + ) + try: + with pytest.raises(ValueError, match="App ID must be provided"): + await client_without_app_id.get_sign_in_resource("connection", activity) + + client_with_app_id = UserTokenClient( + "https://token.example", token="", app_id="app-id" + ) + try: + with pytest.raises(ValueError, match="Activity must have a channel_id"): + await client_with_app_id.get_token_or_sign_in_resource( + "connection", activity + ) + finally: + await client_with_app_id.close() + finally: + await client_without_app_id.close() + + @pytest.mark.asyncio + async def test_missing_user_token_returns_an_empty_token_response(self): + async def handler(request): + return web.Response(status=404) + + app = web.Application() + app.router.add_get("/api/usertoken/GetToken", handler) + server = TestServer(app) + await server.start_server() + client = UserTokenClient(str(server.make_url("/")), token="", app_id="app-id") + try: + token = await client.get_user_token("user-1", "connection", "msteams") + finally: + await client.close() + await server.close() + + assert token.token is None + assert not token + + @pytest.mark.asyncio + async def test_token_service_failure_raises_client_response_error(self): + async def handler(request): + return web.Response(status=500) + + app = web.Application() + app.router.add_get("/api/usertoken/GetTokenStatus", handler) + server = TestServer(app) + await server.start_server() + client = UserTokenClient(str(server.make_url("/")), token="", app_id="app-id") + try: + with pytest.raises(ClientResponseError) as exc_info: + await client.get_token_status("user-1", "msteams") + finally: + await client.close() + await server.close() + + assert exc_info.value.status == 500