Support managed_identity_client_id for DefaultAzureCredential (#39)

This commit is contained in:
you-n-g
2024-07-01 14:24:47 +08:00
committed by GitHub
parent b796b90cce
commit 47dc796b08
2 changed files with 7 additions and 1 deletions
+2
View File
@@ -15,8 +15,10 @@ from pydantic_settings import BaseSettings
class RDAgentSettings(BaseSettings):
# TODO: (xiao) I think most of the config should be in oai.config
use_azure: bool = True
use_azure_token_provider: bool = False
managed_identity_client_id: str | None = None
max_retry: int = 10
retry_wait_seconds: int = 1
dump_chat_cache: bool = False
+5 -1
View File
@@ -292,6 +292,7 @@ class APIBackend:
else:
self.use_azure = self.cfg.use_azure
self.use_azure_token_provider = self.cfg.use_azure_token_provider
self.managed_identity_client_id = self.cfg.managed_identity_client_id
self.chat_api_key = self.cfg.chat_openai_api_key if chat_api_key is None else chat_api_key
self.chat_model = self.cfg.chat_model if chat_model is None else chat_model
@@ -314,7 +315,10 @@ class APIBackend:
if self.use_azure:
if self.use_azure_token_provider:
credential = DefaultAzureCredential()
dac_kwargs = {}
if self.managed_identity_client_id is not None:
dac_kwargs["managed_identity_client_id"] = self.managed_identity_client_id
credential = DefaultAzureCredential(**dac_kwargs)
token_provider = get_bearer_token_provider(
credential,
"https://cognitiveservices.azure.com/.default",