diff --git a/.env.sample b/.env.sample index 51bbc471..3db899b6 100644 --- a/.env.sample +++ b/.env.sample @@ -1,5 +1,14 @@ -GEMINI_PROJECT_ID= -GEMINI_API_KEY= -GITHUB_TOKEN= -OPENROUTER_API_KEY = -OPENROUTER_MODEL = \ No newline at end of file +# LLM provider and keys +LLM_PROVIDER=gemini +GEMINI_API_KEY=your-gemini-api-key +GEMINI_MODEL=gemini-pro + +# DeepSeek (set LLM_PROVIDER=deepseek) +DEEPSEEK_API_KEY=your-deepseek-api-key +DEEPSEEK_BASE_URL=https://api.deepseek.com/v1 +DEEPSEEK_MODEL=deepseek-chat + +# Kimi (set LLM_PROVIDER=kimi) +KIMI_API_KEY=your-kimi-api-key +KIMI_BASE_URL=https://api.moonshot.cn/v1 +KIMI_MODEL=moonshot-v1-8k diff --git a/utils/call_llm.py b/utils/call_llm.py index 70c9e83a..0828f40e 100644 --- a/utils/call_llm.py +++ b/utils/call_llm.py @@ -1,185 +1,105 @@ -from google import genai import os -import logging import json -import requests -from datetime import datetime - -# Configure logging -log_directory = os.getenv("LOG_DIR", "logs") -os.makedirs(log_directory, exist_ok=True) -log_file = os.path.join( - log_directory, f"llm_calls_{datetime.now().strftime('%Y%m%d')}.log" -) - -# Set up logger -logger = logging.getLogger("llm_logger") -logger.setLevel(logging.INFO) -logger.propagate = False # Prevent propagation to root logger -file_handler = logging.FileHandler(log_file, encoding='utf-8') -file_handler.setFormatter( - logging.Formatter("%(asctime)s - %(levelname)s - %(message)s") -) -logger.addHandler(file_handler) - -# Simple cache configuration -cache_file = "llm_cache.json" +import logging + +logger = logging.getLogger(__name__) + +CACHE_PATH = os.environ.get("CACHE_PATH", os.path.expanduser("~/.code_map_cache.json")) def load_cache(): - try: - with open(cache_file, 'r') as f: + if os.path.exists(CACHE_PATH): + logger.debug("Loading cache from %s", CACHE_PATH) + with open(CACHE_PATH, "r") as f: return json.load(f) - except: - logger.warning(f"Failed to load cache.") return {} def save_cache(cache): - try: - with open(cache_file, 'w') as f: - json.dump(cache, f) - except: - logger.warning(f"Failed to save cache") + logger.debug("Saving cache to %s", CACHE_PATH) + with open(CACHE_PATH, "w") as f: + json.dump(cache, f) def get_llm_provider(): - provider = os.getenv("LLM_PROVIDER") - if not provider and (os.getenv("GEMINI_PROJECT_ID") or os.getenv("GEMINI_API_KEY")): - provider = "GEMINI" - # if necessary, add ANTHROPIC/OPENAI - return provider + return os.environ.get("LLM_PROVIDER", "gemini").lower() def _call_llm_provider(prompt: str) -> str: - """ - Call an LLM provider based on environment variables. - Environment variables: - - LLM_PROVIDER: "OLLAMA" or "XAI" - - _MODEL: Model name (e.g., OLLAMA_MODEL, XAI_MODEL) - - _BASE_URL: Base URL without endpoint (e.g., OLLAMA_BASE_URL, XAI_BASE_URL) - - _API_KEY: API key (e.g., OLLAMA_API_KEY, XAI_API_KEY; optional for providers that don't require it) - The endpoint /v1/chat/completions will be appended to the base URL. - """ - logger.info(f"PROMPT: {prompt}") # log the prompt - - # Read the provider from environment variable - provider = os.environ.get("LLM_PROVIDER") - if not provider: - raise ValueError("LLM_PROVIDER environment variable is required") - - # Construct the names of the other environment variables - model_var = f"{provider}_MODEL" - base_url_var = f"{provider}_BASE_URL" - api_key_var = f"{provider}_API_KEY" - - # Read the provider-specific variables - model = os.environ.get(model_var) - base_url = os.environ.get(base_url_var) - api_key = os.environ.get(api_key_var, "") # API key is optional, default to empty string - - # Validate required variables - if not model: - raise ValueError(f"{model_var} environment variable is required") - if not base_url: - raise ValueError(f"{base_url_var} environment variable is required") - - # Append the endpoint to the base URL - url = f"{base_url.rstrip('/')}/v1/chat/completions" - - # Configure headers and payload based on provider - headers = { - "Content-Type": "application/json", - } - if api_key: # Only add Authorization header if API key is provided - headers["Authorization"] = f"Bearer {api_key}" - - payload = { - "model": model, - "messages": [{"role": "user", "content": prompt}], - "temperature": 0.7, - } - - try: - response = requests.post(url, headers=headers, json=payload) - response_json = response.json() # Log the response - logger.info("RESPONSE:\n%s", json.dumps(response_json, indent=2)) - #logger.info(f"RESPONSE: {response.json()}") - response.raise_for_status() - return response.json()["choices"][0]["message"]["content"] - except requests.exceptions.HTTPError as e: - error_message = f"HTTP error occurred: {e}" - try: - error_details = response.json().get("error", "No additional details") - error_message += f" (Details: {error_details})" - except: - pass - raise Exception(error_message) - except requests.exceptions.ConnectionError: - raise Exception(f"Failed to connect to {provider} API. Check your network connection.") - except requests.exceptions.Timeout: - raise Exception(f"Request to {provider} API timed out.") - except requests.exceptions.RequestException as e: - raise Exception(f"An error occurred while making the request to {provider}: {e}") - except ValueError: - raise Exception(f"Failed to parse response as JSON from {provider}. The server might have returned an invalid response.") - -# By default, we Google Gemini 2.5 pro, as it shows great performance for code understanding -def call_llm(prompt: str, use_cache: bool = True) -> str: - # Log the prompt - logger.info(f"PROMPT: {prompt}") + provider = get_llm_provider() + if provider == "gemini": + return _call_llm_gemini(prompt) + elif provider == "deepseek": + return _call_llm_deepseek(prompt) + elif provider == "kimi": + return _call_llm_kimi(prompt) + else: + raise ValueError(f"Unsupported LLM provider: {provider}") + - # Check cache if enabled +def call_llm(prompt: str, use_cache: bool = True) -> str: if use_cache: - # Load cache from disk cache = load_cache() - # Return from cache if exists if prompt in cache: - logger.info(f"RESPONSE: {cache[prompt]}") + logger.debug("Cache hit for prompt") return cache[prompt] - provider = get_llm_provider() - if provider == "GEMINI": - response_text = _call_llm_gemini(prompt) - else: # generic method using a URL that is OpenAI compatible API (Ollama, ...) - response_text = _call_llm_provider(prompt) - - # Log the response - logger.info(f"RESPONSE: {response_text}") + result = _call_llm_provider(prompt) - # Update cache if enabled if use_cache: - # Load cache again to avoid overwrites cache = load_cache() - # Add to cache and save - cache[prompt] = response_text + cache[prompt] = result save_cache(cache) - return response_text + return result def _call_llm_gemini(prompt: str) -> str: - if os.getenv("GEMINI_PROJECT_ID"): - client = genai.Client( - vertexai=True, - project=os.getenv("GEMINI_PROJECT_ID"), - location=os.getenv("GEMINI_LOCATION", "us-central1") - ) - elif os.getenv("GEMINI_API_KEY"): - client = genai.Client(api_key=os.getenv("GEMINI_API_KEY")) - else: - raise ValueError("Either GEMINI_PROJECT_ID or GEMINI_API_KEY must be set in the environment") - model = os.getenv("GEMINI_MODEL", "gemini-2.5-pro-exp-03-25") - response = client.models.generate_content( + import google.generativeai as genai + + api_key = os.environ.get("GEMINI_API_KEY") + if not api_key: + logger.error("GEMINI_API_KEY environment variable not set.") + raise ValueError("GEMINI_API_KEY environment variable not set.") + genai.configure(api_key=api_key) + model_name = os.environ.get("GEMINI_MODEL", "gemini-pro") + logger.debug("Using Gemini model: %s", model_name) + model = genai.GenerativeModel(model_name) + response = model.generate_content(prompt) + return response.text + + +def _call_llm_deepseek(prompt: str) -> str: + from openai import OpenAI + + api_key = os.environ.get("DEEPSEEK_API_KEY") + if not api_key: + logger.error("DEEPSEEK_API_KEY environment variable not set.") + raise ValueError("DEEPSEEK_API_KEY environment variable not set.") + base_url = os.environ.get("DEEPSEEK_BASE_URL", "https://api.deepseek.com/v1") + model = os.environ.get("DEEPSEEK_MODEL", "deepseek-chat") + logger.debug("Using DeepSeek model: %s at %s", model, base_url) + client = OpenAI(api_key=api_key, base_url=base_url) + response = client.chat.completions.create( model=model, - contents=[prompt] + messages=[{"role": "user", "content": prompt}], ) - return response.text + return response.choices[0].message.content + -if __name__ == "__main__": - test_prompt = "Hello, how are you?" +def _call_llm_kimi(prompt: str) -> str: + from openai import OpenAI - # First call - should hit the API - print("Making call...") - response1 = call_llm(test_prompt, use_cache=False) - print(f"Response: {response1}") + api_key = os.environ.get("KIMI_API_KEY") + if not api_key: + logger.error("KIMI_API_KEY environment variable not set.") + raise ValueError("KIMI_API_KEY environment variable not set.") + base_url = os.environ.get("KIMI_BASE_URL", "https://api.moonshot.cn/v1") + model = os.environ.get("KIMI_MODEL", "moonshot-v1-8k") + logger.debug("Using Kimi model: %s at %s", model, base_url) + client = OpenAI(api_key=api_key, base_url=base_url) + response = client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": prompt}], + ) + return response.choices[0].message.content