Skip to content
Closed
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
52 changes: 23 additions & 29 deletions ai/ai_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,8 @@
Generic client that delegates to specific providers (Gemini, OpenAI, Claude, OpenRouter)
"""

from typing import Optional, Dict, Any
from typing import Any, Dict, Optional

from ai.providers import get_provider
from utils.logger import get_logger

Expand All @@ -13,85 +14,78 @@ class AIClient:
Unified AI client that works with multiple providers
Delegates all operations to the configured provider
"""

def __init__(self, config: Dict[str, Any]):
"""
Initialize AI client with specified provider

Args:
config: Configuration dictionary
"""
self.config = config
self.logger = get_logger(config)

# Load the appropriate provider based on config
self.provider = get_provider(config)

# Expose provider's model name
self.model_name = self.provider.get_model_name()
self.logger.info(f"AIClient initialized with {self.provider.__class__.__name__}: {self.model_name}")

self.logger.info(
f"AIClient initialized with {self.provider.__class__.__name__}: {self.model_name}"
)

async def generate(
self,
prompt: str,
system_prompt: Optional[str] = None,
context: Optional[list] = None
self, prompt: str, system_prompt: Optional[str] = None, context: Optional[list] = None
) -> str:
"""
Generate AI response asynchronously

Args:
prompt: User prompt
system_prompt: Optional system instruction
context: Optional conversation history

Returns:
Generated response text
"""
return await self.provider.generate(prompt, system_prompt, context)

def generate_sync(
self,
prompt: str,
system_prompt: Optional[str] = None,
context: Optional[list] = None
self, prompt: str, system_prompt: Optional[str] = None, context: Optional[list] = None
) -> str:
"""
Generate AI response synchronously

Args:
prompt: User prompt
system_prompt: Optional system instruction
context: Optional conversation history

Returns:
Generated response text
"""
return self.provider.generate_sync(prompt, system_prompt, context)

async def generate_with_reasoning(
self,
prompt: str,
system_prompt: str,
task_context: Optional[str] = None
self, prompt: str, system_prompt: str, task_context: Optional[str] = None
) -> Dict[str, Any]:
"""
Generate response with reasoning steps

Args:
prompt: User prompt
system_prompt: System instruction
task_context: Optional task context

Returns:
Dictionary with response and reasoning
"""
return await self.provider.generate_with_reasoning(prompt, system_prompt, task_context)

def get_model_name(self) -> str:
"""Get the current model name"""
return self.model_name

def is_available(self) -> bool:
"""Check if the provider is properly configured"""
return self.provider.is_available()
Expand Down
30 changes: 15 additions & 15 deletions ai/prompt_templates/__init__.py
Original file line number Diff line number Diff line change
@@ -1,27 +1,27 @@
"""Prompt templates package"""

from .planner import (
PLANNER_SYSTEM_PROMPT,
PLANNER_DECISION_PROMPT,
PLANNER_ANALYSIS_PROMPT,
)
from .tool_selector import (
TOOL_SELECTOR_SYSTEM_PROMPT,
TOOL_SELECTION_PROMPT,
TOOL_PARAMETERS_PROMPT,
)
from .analyst import (
ANALYST_SYSTEM_PROMPT,
ANALYST_INTERPRET_PROMPT,
ANALYST_CORRELATION_PROMPT,
ANALYST_FALSE_POSITIVE_PROMPT,
ANALYST_INTERPRET_PROMPT,
ANALYST_SYSTEM_PROMPT,
)
from .planner import (
PLANNER_ANALYSIS_PROMPT,
PLANNER_DECISION_PROMPT,
PLANNER_SYSTEM_PROMPT,
)
from .reporter import (
REPORTER_SYSTEM_PROMPT,
REPORTER_AI_TRACE_PROMPT,
REPORTER_EXECUTIVE_SUMMARY_PROMPT,
REPORTER_TECHNICAL_FINDINGS_PROMPT,
REPORTER_REMEDIATION_PROMPT,
REPORTER_AI_TRACE_PROMPT,
REPORTER_SYSTEM_PROMPT,
REPORTER_TECHNICAL_FINDINGS_PROMPT,
)
from .tool_selector import (
TOOL_PARAMETERS_PROMPT,
TOOL_SELECTION_PROMPT,
TOOL_SELECTOR_SYSTEM_PROMPT,
)

__all__ = [
Expand Down
2 changes: 1 addition & 1 deletion ai/prompt_templates/reporter.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
"""
Prompt templates for the Reporter Agent
Prompt templates for the Reporter Agent
Generates structured penetration testing reports
"""

Expand Down
28 changes: 13 additions & 15 deletions ai/providers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,9 @@
Provider factory and registry for different AI providers
"""

from typing import Dict, Any
from utils.logger import get_logger
from typing import Any, Dict

from utils.logger import get_logger

# Provider registry
PROVIDERS = {
Expand All @@ -20,45 +20,43 @@
def get_provider(config: Dict[str, Any]):
"""
Factory function to create appropriate AI provider

Args:
config: Configuration dictionary

Returns:
Initialized provider instance

Raises:
ValueError: If provider is unknown
RuntimeError: If provider initialization fails
"""
logger = get_logger(config)

# Get selected provider from config
ai_config = config.get("ai", {})
provider_name = ai_config.get("provider", "gemini").lower()

if provider_name not in PROVIDERS:
available = ", ".join(PROVIDERS.keys())
raise ValueError(
f"Unknown provider: {provider_name}. "
f"Available providers: {available}"
)

raise ValueError(f"Unknown provider: {provider_name}. " f"Available providers: {available}")

# Import and instantiate provider
provider_path = PROVIDERS[provider_name]
module_path, class_name = provider_path.rsplit(".", 1)

try:
# Dynamic import
import importlib

module = importlib.import_module(module_path)
provider_class = getattr(module, class_name)

# Create and return provider instance
provider = provider_class(config, logger)
logger.info(f"Loaded provider: {provider_name}")
return provider

except ImportError as e:
raise RuntimeError(
f"Failed to import provider {provider_name}: {e}. "
Expand Down
49 changes: 21 additions & 28 deletions ai/providers/base_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,72 +3,65 @@
Defines the common interface that all AI providers must implement
"""

from abc import ABC, abstractmethod
from typing import Optional, Dict, Any, List
import time
import asyncio
import time
from abc import ABC, abstractmethod
from typing import Any, Dict, Optional


class BaseProvider(ABC):
"""Abstract base class for AI providers"""

def __init__(self, config: Dict[str, Any], logger):
self.config = config
self.logger = logger

# Rate limiting
ai_config = config.get("ai", {})
self.rate_limit = ai_config.get("rate_limit", 60)
self._min_request_interval = 60.0 / self.rate_limit if self.rate_limit > 0 else 0
self._last_request_time = 0.0

@abstractmethod
def _initialize(self):
"""Initialize the provider backend"""
pass

@abstractmethod
async def generate(
self,
prompt: str,
system_prompt: Optional[str] = None,
context: Optional[list] = None
self, prompt: str, system_prompt: Optional[str] = None, context: Optional[list] = None
) -> str:
"""Generate response asynchronously"""
pass

@abstractmethod
def generate_sync(
self,
prompt: str,
system_prompt: Optional[str] = None,
context: Optional[list] = None
self, prompt: str, system_prompt: Optional[str] = None, context: Optional[list] = None
) -> str:
"""Generate response synchronously"""
pass

async def generate_with_reasoning(
self,
prompt: str,
system_prompt: str,
task_context: Optional[str] = None
self, prompt: str, system_prompt: str, task_context: Optional[str] = None
) -> Dict[str, Any]:
"""
Generate response with reasoning steps
Default implementation - can be overridden by providers
"""
full_prompt = f"{prompt}\n\nPlease think through this step-by-step and provide your reasoning."

full_prompt = (
f"{prompt}\n\nPlease think through this step-by-step and provide your reasoning."
)

if task_context:
full_prompt = f"Context: {task_context}\n\n{full_prompt}"

response = await self.generate(full_prompt, system_prompt)

return {
"response": response,
"reasoning": "Provider does not support explicit reasoning extraction"
"reasoning": "Provider does not support explicit reasoning extraction",
}

async def _apply_rate_limit(self):
"""Apply rate limiting between API calls"""
if self._min_request_interval > 0:
Expand All @@ -77,7 +70,7 @@ async def _apply_rate_limit(self):
wait_time = self._min_request_interval - elapsed
await asyncio.sleep(wait_time)
self._last_request_time = time.time()

def _apply_rate_limit_sync(self):
"""Apply rate limiting between API calls (synchronous)"""
if self._min_request_interval > 0:
Expand Down
Loading
Loading