File size: 3,435 Bytes
33516f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
"""React Agent Utils - Helper functions for model creation and tool management.

Purpose:
    Utility functions for creating chat models, resolving API keys, and building tool legends.
"""

import os
from typing import Optional, Dict, Any, List
from dotenv import load_dotenv
from src.react_agent.configuration import Configuration
from src.llm.factory import create_chat_model as _create_chat_model_shared

# Load environment variables from .env file
load_dotenv()


def _resolve_registry_entry(model_name: str, config: Any) -> Dict[str, str]:
    """Resolve model registry entry for a given model name.
    
    Args:
        model_name: Name of the model
        config: Configuration object or RunnableConfig
        
    Returns:
        Dictionary with provider and env_var keys, or empty dict if not found
    """
    cfg = Configuration.from_runnable_config(config)
    registry = getattr(cfg, "model_registry", {}) or {}
    return registry.get(model_name, {})


def get_api_key_for_model(model_name: str, config: Any) -> Optional[str]:
    """Get API key for a model from environment variables."""
    entry = _resolve_registry_entry(model_name, config)
    env_var = entry.get("env_var")
    return os.getenv(env_var) if env_var else None


def get_model_provider(model_name: str) -> Optional[str]:
    """Infer provider from model name."""
    name = (model_name or "").lower()
    if "gemini" in name:
        return "google_genai"
    if "claude" in name:
        return "anthropic"
    if "fireworks" in name or name.startswith("fireworks-ai/"):
        return "fireworks"
    if name.startswith("gpt") or name.startswith("o1"):
        return "openai"
    return None


def create_chat_model(
    model_name: str,
    config: Any = None,
    max_tokens: Optional[int] = None,
    temperature: Optional[float] = None,
    timeout: Optional[int] = None,
) -> Any:
    """Create chat model instance based on model name.

    This delegates to the shared `src.llm.factory` so pipelines + agents share one
    provider selection and env var validation implementation.
    """
    cfg = Configuration.from_runnable_config(config)
    registry_entry = _resolve_registry_entry(model_name, config)
    
    # Enable LangSmith tracing if available
    os.environ.setdefault("LANGCHAIN_TRACING_V2", "true")
    os.environ.setdefault("LANGCHAIN_PROJECT", "react_agent")
    
    max_tokens = max_tokens if max_tokens is not None else cfg.max_tokens
    temperature = temperature if temperature is not None else cfg.temperature
    timeout = timeout if timeout is not None else cfg.timeout

    # Pass the agent's registry through so custom additions continue to work.
    return _create_chat_model_shared(
        model_name=model_name,
        registry=cfg.model_registry,
        temperature=temperature,
        max_tokens=max_tokens,
        timeout=timeout,
    )


def build_tool_legend(tools: List[Any]) -> str:
    """Return a concise legend of available tools."""
    if not tools:
        return "Tools available: none."
    entries = []
    for tool in tools:
        name = getattr(tool, "name", None) or tool.get("name", "unknown")  # type: ignore[attr-defined]
        desc = getattr(tool, "description", "") or getattr(tool, "__doc__", "") or ""
        desc = desc.strip().replace("\n", " ")
        entries.append(f"- {name}: {desc}" if desc else f"- {name}")
    return "Tools available:\n" + "\n".join(entries)