mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 02:43:49 +08:00
185 lines
7.2 KiB
Python
185 lines
7.2 KiB
Python
"""Compatibility gateway for the legacy tracing provider implementations."""
|
|
|
|
import logging
|
|
from typing import Any, override
|
|
|
|
from pydantic import ValidationError
|
|
|
|
from core.ops.entities.config_entity import BaseTracingConfig
|
|
from core.ops.ops_trace_manager import OpsTraceManager, TracingProviderConfigEntry, provider_config_map
|
|
from services.app_tracing_config_service import (
|
|
AppTracingConfigInvalidConfigurationError,
|
|
AppTracingConfigInvalidProviderError,
|
|
AppTracingConfigProcessingError,
|
|
AppTracingConfigVerificationFailedError,
|
|
TracingConfigProviderGateway,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_PROJECT_URL_FALLBACKS = {
|
|
"arize": "https://app.arize.com/",
|
|
"phoenix": "https://app.phoenix.arize.com/projects/",
|
|
"langsmith": "https://smith.langchain.com/",
|
|
"opik": "https://www.comet.com/opik/",
|
|
"weave": "https://wandb.ai/",
|
|
"aliyun": "https://arms.console.aliyun.com/",
|
|
"tencent": "https://console.cloud.tencent.com/apm",
|
|
"mlflow": "http://localhost:5000/",
|
|
"databricks": "https://www.databricks.com/",
|
|
}
|
|
|
|
|
|
class OpsTraceManagerGateway(TracingConfigProviderGateway):
|
|
@override
|
|
def validate_provider(self, tracing_provider: str) -> None:
|
|
self._provider_config(tracing_provider)
|
|
|
|
@override
|
|
def prepare_new_config(
|
|
self,
|
|
*,
|
|
workspace_id: str,
|
|
tracing_provider: str,
|
|
tracing_config: dict[str, Any],
|
|
) -> dict[str, Any]:
|
|
provider_config = self._provider_config(tracing_provider)
|
|
normalized_config = self._normalize_config(provider_config, tracing_config)
|
|
self._verify_config(normalized_config, tracing_provider)
|
|
|
|
project_url = self._get_project_url_for_create(normalized_config, tracing_provider)
|
|
try:
|
|
encrypted_config = OpsTraceManager.encrypt_tracing_config(
|
|
workspace_id,
|
|
tracing_provider,
|
|
normalized_config,
|
|
)
|
|
except Exception as error:
|
|
raise AppTracingConfigProcessingError from error
|
|
if project_url:
|
|
encrypted_config["project_url"] = project_url
|
|
return encrypted_config
|
|
|
|
@override
|
|
def prepare_updated_config(
|
|
self,
|
|
*,
|
|
workspace_id: str,
|
|
tracing_provider: str,
|
|
tracing_config: dict[str, Any],
|
|
current_tracing_config: dict[str, Any] | None,
|
|
) -> dict[str, Any]:
|
|
provider_config = self._provider_config(tracing_provider)
|
|
self._validate_config(provider_config, tracing_config)
|
|
try:
|
|
encrypted_config = OpsTraceManager.encrypt_tracing_config(
|
|
workspace_id,
|
|
tracing_provider,
|
|
dict(tracing_config),
|
|
current_tracing_config,
|
|
)
|
|
decrypted_config = OpsTraceManager.decrypt_tracing_config(
|
|
workspace_id,
|
|
tracing_provider,
|
|
encrypted_config,
|
|
)
|
|
except Exception as error:
|
|
raise AppTracingConfigProcessingError from error
|
|
|
|
self._verify_config(decrypted_config, tracing_provider)
|
|
return encrypted_config
|
|
|
|
@override
|
|
def present_config(
|
|
self,
|
|
*,
|
|
workspace_id: str,
|
|
tracing_provider: str,
|
|
tracing_config: dict[str, Any] | None,
|
|
) -> dict[str, Any]:
|
|
if tracing_config is None:
|
|
raise AppTracingConfigProcessingError
|
|
|
|
try:
|
|
decrypted_config = OpsTraceManager.decrypt_tracing_config(
|
|
workspace_id,
|
|
tracing_provider,
|
|
tracing_config,
|
|
)
|
|
presented_config = OpsTraceManager.obfuscated_decrypt_token(tracing_provider, decrypted_config)
|
|
except Exception as error:
|
|
raise AppTracingConfigProcessingError from error
|
|
|
|
if tracing_provider == "langfuse" and not decrypted_config.get("project_key"):
|
|
try:
|
|
project_key = OpsTraceManager.get_trace_config_project_key(decrypted_config, tracing_provider)
|
|
presented_config["project_url"] = f"{decrypted_config.get('host')}/project/{project_key}"
|
|
except Exception:
|
|
presented_config["project_url"] = f"{decrypted_config.get('host')}/"
|
|
elif tracing_provider in _PROJECT_URL_FALLBACKS and not decrypted_config.get("project_url"):
|
|
try:
|
|
presented_config["project_url"] = OpsTraceManager.get_trace_config_project_url(
|
|
decrypted_config,
|
|
tracing_provider,
|
|
)
|
|
except Exception:
|
|
presented_config["project_url"] = _PROJECT_URL_FALLBACKS[tracing_provider]
|
|
|
|
return presented_config
|
|
|
|
@staticmethod
|
|
def _provider_config(tracing_provider: str) -> TracingProviderConfigEntry:
|
|
try:
|
|
return provider_config_map[tracing_provider]
|
|
except KeyError as error:
|
|
raise AppTracingConfigInvalidProviderError(tracing_provider) from error
|
|
|
|
@staticmethod
|
|
def _normalize_config(
|
|
provider_config: TracingProviderConfigEntry,
|
|
tracing_config: dict[str, Any],
|
|
) -> dict[str, Any]:
|
|
normalized_config = dict(tracing_config)
|
|
default_config = OpsTraceManagerGateway._validate_config(provider_config, normalized_config)
|
|
|
|
default_values = default_config.model_dump()
|
|
for key in provider_config["other_keys"]:
|
|
if normalized_config.get(key) == "":
|
|
normalized_config[key] = default_values.get(key)
|
|
return normalized_config
|
|
|
|
@staticmethod
|
|
def _validate_config(
|
|
provider_config: TracingProviderConfigEntry,
|
|
tracing_config: dict[str, Any],
|
|
) -> BaseTracingConfig:
|
|
config_class: type[BaseTracingConfig] = provider_config["config_class"]
|
|
try:
|
|
return config_class.model_validate(tracing_config)
|
|
except ValidationError as error:
|
|
raise AppTracingConfigInvalidConfigurationError from error
|
|
|
|
@staticmethod
|
|
def _verify_config(tracing_config: dict[str, Any], tracing_provider: str) -> None:
|
|
try:
|
|
is_effective = OpsTraceManager.check_trace_config_is_effective(tracing_config, tracing_provider)
|
|
except ValueError as error:
|
|
logger.warning("Tracing configuration verification failed for provider %s", tracing_provider, exc_info=True)
|
|
raise AppTracingConfigVerificationFailedError from error
|
|
if not is_effective:
|
|
raise AppTracingConfigVerificationFailedError
|
|
|
|
@staticmethod
|
|
def _get_project_url_for_create(tracing_config: dict[str, Any], tracing_provider: str) -> str | None:
|
|
try:
|
|
if tracing_provider in ("arize", "phoenix"):
|
|
return OpsTraceManager.get_trace_config_project_url(tracing_config, tracing_provider)
|
|
if tracing_provider == "langfuse":
|
|
project_key = OpsTraceManager.get_trace_config_project_key(tracing_config, tracing_provider)
|
|
return f"{tracing_config.get('host')}/project/{project_key}"
|
|
if tracing_provider in ("langsmith", "opik", "mlflow", "databricks", "tencent"):
|
|
return OpsTraceManager.get_trace_config_project_url(tracing_config, tracing_provider)
|
|
except Exception:
|
|
return None
|
|
return None
|