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
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) Microsoft Corporation.
# Licensed under the MIT License.

from functools import lru_cache
from importlib.metadata import version

import grpc
Expand All @@ -17,6 +18,22 @@
)


@lru_cache(maxsize=1)
def _get_sdk_version() -> str:
"""Return the installed version of the azuremanaged package.

Resolving the version walks distribution metadata on disk, so the result is
cached and shared by every interceptor instance instead of being recomputed
on each client or worker construction. Falls back to ``"unknown"`` when the
version cannot be determined.
"""
try:
return version('durabletask-azuremanaged')
except Exception:
# Fallback if version cannot be determined
return "unknown"


class DTSDefaultClientInterceptorImpl (DefaultClientInterceptorImpl):
"""The class implements a UnaryUnaryClientInterceptor, UnaryStreamClientInterceptor,
StreamUnaryClientInterceptor and StreamStreamClientInterceptor from grpc to add an
Expand All @@ -27,13 +44,7 @@ def __init__(
token_credential: TokenCredential | None,
taskhub_name: str,
worker_id: str | None = None):
try:
# Get the version of the azuremanaged package
sdk_version = version('durabletask-azuremanaged')
except Exception:
# Fallback if version cannot be determined
sdk_version = "unknown"
user_agent = f"durabletask-python/{sdk_version}"
user_agent = f"durabletask-python/{_get_sdk_version()}"
self._metadata = [
("taskhub", taskhub_name),
("x-user-agent", user_agent)] # 'user-agent' is a reserved header; use 'x-user-agent'
Expand Down Expand Up @@ -81,13 +92,7 @@ class DTSAsyncDefaultClientInterceptorImpl(DefaultAsyncClientInterceptorImpl):
(task hub name, user agent, and authentication token) to all async calls."""

def __init__(self, token_credential: AsyncTokenCredential | None, taskhub_name: str):
try:
# Get the version of the azuremanaged package
sdk_version = version('durabletask-azuremanaged')
except Exception:
# Fallback if version cannot be determined
sdk_version = "unknown"
user_agent = f"durabletask-python/{sdk_version}"
user_agent = f"durabletask-python/{_get_sdk_version()}"
self._metadata = [
("taskhub", taskhub_name),
("x-user-agent", user_agent)]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,19 +5,25 @@
import unittest
from concurrent import futures
from datetime import datetime, timedelta, timezone
from importlib.metadata import version
from importlib.metadata import PackageNotFoundError, version
import threading
import time
from typing import Any
from unittest import mock

import grpc
from azure.core.credentials import AccessToken

from durabletask.azuremanaged.client import DurableTaskSchedulerClient
from durabletask.azuremanaged.internal import durabletask_grpc_interceptor
from durabletask.azuremanaged.internal.access_token_manager import (
AccessTokenManager,
AsyncAccessTokenManager,
)
from durabletask.azuremanaged.internal.durabletask_grpc_interceptor import (
DTSAsyncDefaultClientInterceptorImpl,
DTSDefaultClientInterceptorImpl,
)
from durabletask.azuremanaged.worker import DurableTaskSchedulerWorker
from durabletask.internal.grpc_interceptor import DefaultClientInterceptorImpl
from durabletask.internal import orchestrator_service_pb2 as pb
Expand Down Expand Up @@ -197,6 +203,58 @@ def test_worker_construction_does_not_acquire_token(self):
self.assertEqual(0, credential.calls, "Constructing a worker must not acquire a token")


class TestSdkVersionCaching(unittest.TestCase):
"""Tests that the azuremanaged SDK version is resolved once and reused."""

def setUp(self):
durabletask_grpc_interceptor._get_sdk_version.cache_clear()

def tearDown(self):
# Drop any patched value so later tests observe the real package version.
durabletask_grpc_interceptor._get_sdk_version.cache_clear()

def test_version_resolved_once_across_interceptor_constructions(self):
"""The distribution metadata lookup happens once, not per interceptor."""
with mock.patch.object(
durabletask_grpc_interceptor, "version", return_value="1.2.3") as mock_version:
Comment thread
andystaples marked this conversation as resolved.
client_interceptor = DTSDefaultClientInterceptorImpl(None, "test-taskhub")
worker_interceptor = DTSDefaultClientInterceptorImpl(
None, "test-taskhub", worker_id="test-worker-id")
async_interceptor = DTSAsyncDefaultClientInterceptorImpl(None, "test-taskhub")

mock_version.assert_called_once_with('durabletask-azuremanaged')

# The cached value is reused verbatim, and the metadata keys and their
# order are unchanged.
self.assertEqual(
[("taskhub", "test-taskhub"), ("x-user-agent", "durabletask-python/1.2.3")],
client_interceptor._metadata)
self.assertEqual(
[("taskhub", "test-taskhub"),
("x-user-agent", "durabletask-python/1.2.3"),
("workerid", "test-worker-id")],
worker_interceptor._metadata)
self.assertEqual(
[("taskhub", "test-taskhub"), ("x-user-agent", "durabletask-python/1.2.3")],
async_interceptor._metadata)

def test_unknown_fallback_when_version_cannot_be_determined(self):
"""A missing distribution still yields the 'unknown' user agent fallback."""
with mock.patch.object(
durabletask_grpc_interceptor,
"version",
side_effect=PackageNotFoundError('durabletask-azuremanaged')) as mock_version:
client_interceptor = DTSDefaultClientInterceptorImpl(None, "test-taskhub")
async_interceptor = DTSAsyncDefaultClientInterceptorImpl(None, "test-taskhub")

# The failed lookup is cached too, so it is not retried per interceptor.
mock_version.assert_called_once_with('durabletask-azuremanaged')
self.assertEqual(
"durabletask-python/unknown", dict(client_interceptor._metadata)["x-user-agent"])
self.assertEqual(
"durabletask-python/unknown", dict(async_interceptor._metadata)["x-user-agent"])


class TestAccessTokenManagerThreadSafety(unittest.TestCase):

@staticmethod
Expand Down
Loading