diff --git a/plugins/platforms/teams/adapter.py b/plugins/platforms/teams/adapter.py index 55aa030310..f6b357208f 100644 --- a/plugins/platforms/teams/adapter.py +++ b/plugins/platforms/teams/adapter.py @@ -514,20 +514,28 @@ _ALLOWED_TEAMS_SERVICE_HOSTS = frozenset({ }) -def _is_botframework_attachment_host(host: str) -> bool: - """True if ``host`` (lowercased) is a Bot Framework connector host. +def _is_botframework_attachment_url(url: str) -> bool: + """True if ``url`` points at a Bot Framework connector attachment host. - Dot-anchored suffix match so attacker hosts like - ``evil-trafficmanager.net`` / ``notbotframework.com`` never receive the - bot's bearer token (same threat model as _ALLOWED_TEAMS_SERVICE_HOSTS: - the token must only ever be sent to Bot Framework infrastructure). + Exact-match against ``_ALLOWED_TEAMS_SERVICE_HOSTS`` — the same allowlist + that gates where outbound sends may carry a freshly minted bearer token — + plus scheme/port sanity: only https on the default port qualifies. A + lookalike host must never receive the bot's bearer token: note that any + Azure customer can register ``.trafficmanager.net`` Traffic Manager + profiles, so a suffix match would not be safe either. New Bot Framework + regions are allowlist additions, not predicate changes. """ - return ( - host == "trafficmanager.net" - or host.endswith(".trafficmanager.net") - or host == "botframework.com" - or host.endswith(".botframework.com") - ) + try: + from urllib.parse import urlparse + + parsed = urlparse(url) + if parsed.scheme != "https": + return False + if parsed.port not in (None, 443): + return False + return parsed.hostname in _ALLOWED_TEAMS_SERVICE_HOSTS + except Exception: + return False # Conservative pattern for Bot Framework conversation IDs. Real values # combine digits, colons, hyphens, dots, '@', and the ``thread.skype`` / @@ -796,8 +804,11 @@ class TeamsAdapter(BasePlatformAdapter): self._client_id = extra.get("client_id") or os.getenv("TEAMS_CLIENT_ID", "") self._client_secret = extra.get("client_secret") or _get_scoped_secret("TEAMS_CLIENT_SECRET", "") self._tenant_id = extra.get("tenant_id") or os.getenv("TEAMS_TENANT_ID", "") - # (token, expiry unix ts) for Bot Framework connector attachment auth + # (token, expiry monotonic ts) for Bot Framework connector attachment + # auth; refreshed under _bf_token_lock so concurrent attachments + # can't stampede the token endpoint. self._bf_token_cache: Optional[tuple] = None + self._bf_token_lock: Optional[asyncio.Lock] = None self._port = _coerce_port( extra.get("port") or os.getenv("TEAMS_PORT", str(_DEFAULT_PORT)) ) @@ -922,58 +933,67 @@ class TeamsAdapter(BasePlatformAdapter): Needed to download connector attachments (smba.trafficmanager.net /v3/attachments/...), which -- unlike SharePoint file downloadUrls -- are NOT pre-authenticated and return 401 without the bot's own - token. Token is cached until ~5 minutes before expiry. + token. Token is cached until ~5 minutes before expiry. The refresh + is serialized by an asyncio lock (lazily created on first use — + ``asyncio.Lock()`` at __init__ time would bind to the wrong event + loop on Python < 3.10) so concurrent attachments share one POST. """ import time import httpx - cached = self._bf_token_cache - if cached and cached[1] > time.time() + 300: - return cached[0] + # The gateway may run adapters on a loop created after __init__; + # bind the lock on first use instead of at construction. + lock = self._bf_token_lock + if lock is None: + lock = self._bf_token_lock = asyncio.Lock() + async with lock: + cached = self._bf_token_cache + if cached and cached[1] > time.monotonic() + 300: + return cached[0] - client_id = self._client_id - client_secret = self._client_secret - tenant_id = self._tenant_id - if not (client_id and client_secret and tenant_id): - raise ValueError("Missing TEAMS_CLIENT_ID/SECRET/TENANT_ID for attachment auth") + client_id = self._client_id + client_secret = self._client_secret + tenant_id = self._tenant_id + if not (client_id and client_secret and tenant_id): + raise ValueError("Missing TEAMS_CLIENT_ID/SECRET/TENANT_ID for attachment auth") - async with httpx.AsyncClient(timeout=15.0) as client: - resp = await client.post( - f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token", - data={ - "grant_type": "client_credentials", - "client_id": client_id, - "client_secret": client_secret, - "scope": "https://api.botframework.com/.default", - }, - ) - resp.raise_for_status() - payload = resp.json() - token = payload["access_token"] - self._bf_token_cache = (token, time.time() + int(payload.get("expires_in", 3600))) - return token + async with httpx.AsyncClient(timeout=15.0) as client: + resp = await client.post( + f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token", + data={ + "grant_type": "client_credentials", + "client_id": client_id, + "client_secret": client_secret, + "scope": "https://api.botframework.com/.default", + }, + ) + resp.raise_for_status() + payload = resp.json() + token = payload["access_token"] + expires_in = float(payload.get("expires_in", 3600) or 3600) + self._bf_token_cache = (token, time.monotonic() + expires_in) + return token async def _fetch_attachment_bytes(self, url: str, timeout: float = 30.0) -> bytes: """Download attachment bytes with SSRF protection. Teams file attachments carry pre-authenticated SharePoint download URLs (no extra auth header needed). Bot Framework connector - attachment URLs (pasted/inline images on smba.trafficmanager.net / - botframework.com hosts) require the bot's bearer token -- detected - below and fetched with auth. Validates the URL against the SSRF - guard and follows redirects through the shared redirect guard, - matching the cache_*_from_url helpers in gateway.platforms.base. + attachment URLs (pasted/inline images on _ALLOWED_TEAMS_SERVICE_HOSTS + hosts) require the bot's bearer token -- detected below and fetched + with auth. Validates the URL against the SSRF guard, streams the + body through the shared inbound media cap, and follows redirects + through the shared redirect guard, matching the cache_*_from_url + helpers in gateway.platforms.base. """ - from urllib.parse import urlparse from tools.url_safety import create_ssrf_safe_async_client, is_safe_url - from gateway.platforms.base import _ssrf_redirect_guard + from gateway.platforms.base import _ssrf_redirect_guard, _read_httpx_body_with_limit if not is_safe_url(url): raise ValueError("Blocked unsafe attachment URL (SSRF protection)") headers = {"User-Agent": "Mozilla/5.0 (compatible; HermesAgent/1.0)"} - host = (urlparse(url).hostname or "").lower() - if _is_botframework_attachment_host(host): + if _is_botframework_attachment_url(url): try: headers["Authorization"] = f"Bearer {await self._get_botframework_token()}" except Exception as e: @@ -984,9 +1004,12 @@ class TeamsAdapter(BasePlatformAdapter): follow_redirects=True, event_hooks={"response": [_ssrf_redirect_guard]}, ) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - return response.content + async with client.stream("GET", url, headers=headers) as response: + response.raise_for_status() + # Stream through the shared inbound media cap (matches + # cache_image_from_url) instead of buffering .content — a + # lying Content-Length must not OOM the gateway. + return await _read_httpx_body_with_limit(response, media_type="attachment") async def _on_message(self, ctx: ActivityContext[MessageActivity]) -> None: """Process an incoming Teams message and dispatch to the gateway.""" @@ -1088,9 +1111,7 @@ class TeamsAdapter(BasePlatformAdapter): if content_url and content_type.startswith("image/"): try: - from urllib.parse import urlparse as _urlparse - _host = (_urlparse(content_url).hostname or "").lower() - if _is_botframework_attachment_host(_host): + if _is_botframework_attachment_url(content_url): # Bot Framework connector URL: needs the bot's own # bearer token; the generic cache helper sends none. data = await self._fetch_attachment_bytes(content_url) diff --git a/tests/gateway/test_teams.py b/tests/gateway/test_teams.py index a167493393..dc3b489e2f 100644 --- a/tests/gateway/test_teams.py +++ b/tests/gateway/test_teams.py @@ -678,21 +678,23 @@ class TestTeamsBotFrameworkAttachments: return att @pytest.mark.anyio - async def test_bf_host_predicate_dot_anchored(self): - """The host check must be dot-anchored: attacker lookalike hosts must - NOT receive the bot's bearer token.""" - f = _teams_mod._is_botframework_attachment_host - assert f("smba.trafficmanager.net") - assert f("emea.smba.trafficmanager.net") - assert f("smba.trafficmanager.net") # exact apex - assert f("api.botframework.com") - assert f("botframework.com") - # Attacker lookalikes — the pre-followup suffix check matched these - assert not f("evil-trafficmanager.net") - assert not f("notbotframework.com") - assert not f("trafficmanager.net.evil.com") + async def test_bf_url_predicate_exact_match_allowlist(self): + """Only exact allowlisted hosts on https default port may receive the + bot's bearer token — lookalikes, other schemes, and non-443 ports must + NOT (any Azure customer can register .trafficmanager.net).""" + f = _teams_mod._is_botframework_attachment_url + assert f("https://smba.trafficmanager.net/emea/v3/attachments/x") + assert f("https://smba.infra.gov.teams.microsoft.us/amer/v3/attachments/x") + assert f("https://smba.trafficmanager.net:443/emea/v3/attachments/x") + # Attacker lookalikes / non-allowlisted / wrong scheme / wrong port + assert not f("https://evil-trafficmanager.net/steal") + assert not f("https://emea.smba.trafficmanager.net/v3/attachments/x") + assert not f("https://notbotframework.com/steal") + assert not f("https://trafficmanager.net.evil.com/steal") + assert not f("http://smba.trafficmanager.net/v3/attachments/x") + assert not f("https://smba.trafficmanager.net:444/v3/attachments/x") assert not f("") - assert not f("sharepoint.com") + assert not f("https://sharepoint.com/x") @pytest.mark.anyio async def test_bf_image_routes_through_authenticated_fetch(self): @@ -745,6 +747,26 @@ class TestTeamsBotFrameworkAttachments: captured = {} + class _FakeStreamResponse: + def __init__(self): + self.headers = {} + + def raise_for_status(self): + pass + + async def aiter_bytes(self): + yield b"\x89PNG fake" + + class _FakeStreamCtx: + def __init__(self, response): + self._response = response + + async def __aenter__(self): + return self._response + + async def __aexit__(self, *a): + return None + class _FakeClient: def __init__(self, **kw): pass @@ -755,30 +777,73 @@ class TestTeamsBotFrameworkAttachments: async def __aexit__(self, *a): return None - async def get(self, url, headers=None): + def stream(self, method, url, headers=None): captured["headers"] = headers or {} - return SimpleNamespace( - status_code=200, - content=b"\x89PNG fake", - raise_for_status=lambda: None, - ) + return _FakeStreamCtx(_FakeStreamResponse()) with patch("tools.url_safety.create_ssrf_safe_async_client", lambda **kw: _FakeClient()), \ patch("tools.url_safety.is_safe_url", lambda url: True): # BF host: bearer attached - await adapter._fetch_attachment_bytes("https://smba.trafficmanager.net/emea/v3/attachments/x") + data = await adapter._fetch_attachment_bytes("https://smba.trafficmanager.net/emea/v3/attachments/x") assert captured["headers"].get("Authorization") == "Bearer the-token" + assert data == b"\x89PNG fake" adapter._get_botframework_token = AsyncMock(return_value="the-token") with patch("tools.url_safety.create_ssrf_safe_async_client", lambda **kw: _FakeClient()), \ patch("tools.url_safety.is_safe_url", lambda url: True): - # Attacker lookalike host: NO bearer (dot-anchored check) + # Attacker lookalike host: NO bearer (exact-match allowlist) await adapter._fetch_attachment_bytes("https://evil-trafficmanager.net/steal") assert "Authorization" not in captured["headers"], ( "bearer token must not be sent to attacker lookalike hosts" ) adapter._get_botframework_token.assert_not_awaited() + @pytest.mark.anyio + async def test_token_refresh_is_serialized_under_lock(self): + """Two concurrent token fetches on a cold cache share ONE POST — + the lock prevents a token-endpoint stampede.""" + import asyncio as _asyncio + + adapter = self._make_adapter() + posts = [] + release = _asyncio.Event() + + class _TokenResp: + status_code = 200 + + def raise_for_status(self): + pass + + def json(self): + return {"access_token": "tok-1", "expires_in": 3600} + + class _SlowTokenClient: + def __init__(self, **kw): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + return None + + async def post(self, url, data=None): + posts.append((url, dict(data or {}))) + await release.wait() # hold both callers at the STS door + return _TokenResp() + + async def release_later(): + await _asyncio.sleep(0.05) + release.set() + + with patch("httpx.AsyncClient", _SlowTokenClient): + t1 = _asyncio.create_task(adapter._get_botframework_token()) + t2 = _asyncio.create_task(adapter._get_botframework_token()) + await release_later() + tok1, tok2 = await t1, await t2 + assert tok1 == "tok-1" and tok2 == "tok-1" + assert len(posts) == 1, f"concurrent cold-cache fetches must share one POST, got {len(posts)}" + @pytest.mark.anyio async def test_token_acquisition_and_cache_reuse(self): adapter = self._make_adapter() @@ -829,6 +894,28 @@ class TestTeamsBotFrameworkAttachments: captured = {} + class _FakeStreamResponse: + def __init__(self): + self.headers = {} + + def raise_for_status(self): + raise _httpx.HTTPStatusError( + "401", request=MagicMock(), response=MagicMock(status_code=401) + ) + + async def aiter_bytes(self): + yield b"" + + class _FakeStreamCtx: + def __init__(self, response): + self._response = response + + async def __aenter__(self): + return self._response + + async def __aexit__(self, *a): + return None + class _FakeClient: def __init__(self, **kw): pass @@ -839,17 +926,9 @@ class TestTeamsBotFrameworkAttachments: async def __aexit__(self, *a): return None - async def get(self, url, headers=None): + def stream(self, method, url, headers=None): captured["headers"] = headers or {} - return SimpleNamespace( - status_code=401, - content=b"", - raise_for_status=lambda: (_ for _ in ()).throw( - _httpx.HTTPStatusError( - "401", request=MagicMock(), response=SimpleNamespace(status_code=401) - ) - ), - ) + return _FakeStreamCtx(_FakeStreamResponse()) with patch("tools.url_safety.create_ssrf_safe_async_client", lambda **kw: _FakeClient()), \ patch("tools.url_safety.is_safe_url", lambda url: True):