diff --git a/CLAUDE.md b/CLAUDE.md index 0efd617b..a63cad8c 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -33,20 +33,22 @@ See `docs/design_philosophy.md` for full design principles. ### Base System (`akd/_base/`) -Everything inherits from `AbstractBase` or `UnrestrictedAbstractBase`. See `docs/specs/AKD_BASE.md` for the full reference. +Everything inherits from `AbstractBase`. See `docs/specs/AKD_BASE.md` for the full reference. **Key exports from `akd._base`:** -| Class | Purpose | +| Class / Protocol | Purpose | |-------|---------| -| `AbstractBase` | Strict base with schema validation (agents, tools) | +| `AbstractBase` | Concrete base with schema validation, streaming, and config binding (agents, tools) | | `InputSchema` / `OutputSchema` / `IOSchema` | Typed schemas with required docstrings | | `BaseConfig` | Configuration base | -| `StreamEvent` / `StreamingMixin` | Streaming event system | -| `ToolCall` / `ToolResult` / `ToolCallingMixin` | Tool calling infrastructure | +| `ConfigBindingMixin` | Opt-in: config property binding + metadata binding (already on `AbstractBase`) | +| `AKDExecutable` / `AKDTool` / `RunContextProtocol` | Structural protocols — framework adapters satisfy these without inheriting `BaseAgent` / `BaseTool` | +| `StreamEvent` | Streaming event hierarchy | +| `ToolCall` / `ToolResult` | Tool calling data models | | `RunContext` | Execution context passed to agents (in `_base/structures.py`) | | `HumanResponse` | Human reply for HITL resumption (in `_base/structures.py`) | -| `Memory` | Message storage with session lifecycle | +| `validate_input` / `validate_output` | Standalone schema validators for non-inheriting adapters | ### Adding Agents or Tools @@ -97,7 +99,6 @@ See specs: **`akd/structures.py`** (public, also re-exports `HumanResponse`): - `SearchResult` / `SearchResultItem` — search results with metadata -- `ExtractionSchema` / `SingleEstimation` — extraction output schemas - `HumanResponse` — re-exported for backend convenience ## Project Structure @@ -107,26 +108,20 @@ akd/ ├── _base/ # Base classes, streaming, tool calling, memory, structures ├── agents/ # Agent implementations │ ├── _base.py # BaseAgent, InstructorBaseAgent, LiteLLMInstructorBaseAgent -│ ├── search/ # SearchAgent, DeepLitSearchAgent, ControlledSearchAgent, AspectSearchAgent -│ ├── gap_analysis/ # GapAgent -│ ├── extraction.py # EstimationExtractionAgent -│ ├── query.py # QueryAgent, FollowUpQueryAgent -│ ├── relevancy.py # Relevancy checking agents -│ ├── storm/ # STORM workflow agent -│ └── intents.py # Intent detection +│ └── relevancy.py # RelevancyAgent, MultiRubricRelevancyAgent ├── tools/ # Tool implementations │ ├── _base.py # BaseTool, BaseToolConfig │ ├── human.py # HumanTool (HITL) -│ ├── search/ # SearxNG, Serper, SemanticScholar, CodeSearch, Composite, Pipeline +│ ├── search/ # SearxNG, Serper, SemanticScholar, Composite, Pipeline │ ├── scrapers/ # Web, PDF, Crawl4AI, PyPaperBot, Docling scrapers │ ├── resolvers/ # DOI, Arxiv, ADS, Unpaywall resolvers │ ├── reranker.py # CrossEncoder, NoOp rerankers │ ├── relevancy.py # Relevancy checker │ └── source_validator.py -├── configs/ # Configuration (project, prompts, lit, storm) +├── configs/ # Configuration (project, prompts, lit) ├── guardrails/ # Safety and validation guardrails -├── mapping/ # Agent/tool registry and field mapping -├── planner/ # Workflow planning +├── mapping/ # Runtime agent registry + field-mapping machinery +├── planner/ # Workflow planning (agents registered at runtime) └── structures.py # Public data structures tests/ # Mirrors akd/ structure — pytest + asyncio + xdist diff --git a/README.md b/README.md index 599b1e62..46440b7d 100644 --- a/README.md +++ b/README.md @@ -79,22 +79,13 @@ This works across any transport — REST APIs, WebSockets, CLI — because the p | Category | Agent | Description | |----------|-------|-------------| -| **Research** | `DeepLitSearchAgent` | Multi-agent deep literature search with triage, clarification, and synthesis | -| | `ControlledSearchAgent` | Controlled search with configurable parameters | -| | `AspectSearchAgent` | Interview-pattern multi-aspect search | -| | `CodeSearchAgent` | Code repository search | -| | `QuestionAnsweringAgent` | QA over retrieved content | -| **Analysis** | `GapAgent` | Research gap identification via knowledge graphs | -| | `EstimationExtractionAgent` | Intent-based data extraction | -| | `StormAgent` | Structured narrative generation | -| **Utility** | `IntentAgent` | User intent classification | -| | `QueryAgent` | Query reformulation and refinement | -| | `FollowUpQueryAgent` | Follow-up query generation | -| | `RelevancyAgent` | Binary relevance classification | +| **Utility** | `RelevancyAgent` | Binary relevance classification | | | `MultiRubricRelevancyAgent` | Multi-dimensional relevance scoring | | **Base** | `BaseAgent` | Core agent with streaming, tool calling, HITL, message trimming | | | `LiteLLMInstructorBaseAgent` | Structured Pydantic output via Instructor | +Domain-specific agents live in downstream packages and can be registered at runtime via `AgentRegistry.register_agent(YourAgent)`. + ## Out-of-Box Tools | Category | Tool | Description | @@ -214,6 +205,27 @@ dependencies = [ ] ``` +**Optional extras:** pull in extra dependencies for specific features. + +| Extra | What it pulls in | Install when you... | +|---|---|---| +| `serializer` | `langgraph` | use `AKDSerializer` as a langgraph checkpoint serde (e.g. `AsyncPostgresSaver(serde=AKDSerializer())`) | +| `ml` | `pandas`, `sentence-transformers`, `docling`, `deepeval` | need ML-backed rerankers, scrapers, or eval tools | +| `dev` | `pytest`, `pytest-asyncio`, `pytest-cov`, `pytest-xdist`, `pre-commit`, `memray`, `scalene` | run the test suite or hack on akd itself | +| `local` | `marimo`, `jupyter`, `ipykernel`, `ipywidgets` | run the marimo notebooks under `notebooks/` | + +```bash +# As a dependency, with an extra: +uv pip install "akd[serializer] @ git+https://github.com/NASA-IMPACT/accelerated-discovery.git@develop" +``` + +```toml +# In your pyproject.toml: +dependencies = [ + "akd[serializer] @ git+https://github.com/NASA-IMPACT/accelerated-discovery.git@develop", +] +``` + **For local development:** ```bash @@ -221,18 +233,24 @@ dependencies = [ uv venv --python 3.12 source .venv/bin/activate -# Install dependencies +# Install core dependencies uv sync -# For development (includes testing tools) +# With development tooling (pytest, pre-commit, profilers) uv sync --extra dev -# For local development (includes marimo and other local tools) +# With notebooks (marimo, jupyter) uv sync --extra dev --extra local -# For ML features (includes sentence-transformers, docling, deepeval) +# With ML extras (pandas, sentence-transformers, docling, deepeval) uv sync --extra ml +# With the langgraph checkpoint serde (AKDSerializer) +uv sync --extra serializer + +# Combine extras freely, e.g. full dev setup: +uv sync --extra dev --extra local --extra ml --extra serializer + # Setup environment variables cp .env.example .env # Edit .env with your API keys @@ -259,13 +277,6 @@ scripts/ # Utility scripts and demos tests/ # Test suite (mirrors akd/ structure) ``` -## Roadmap - -These features are part of the design vision but not yet fully implemented: - -- **Conflict Agent** — a dedicated agent that specifically searches for contradictory evidence and conflicting findings across sources. Currently, conflict detection is a design principle (see [Design Philosophy](docs/design_philosophy.md)) but lacks a standalone agent implementation. -- **Full Attribution Chain** — end-to-end traceability from final claims back to specific source sentences. Partial support exists today: `GapAgent` provides `attributed_source_answers` and the guardrail system includes an `ATTRIBUTION` risk category. - ## Contributing See [CONTRIBUTING.md](CONTRIBUTING.md) for setup, style guide, branch conventions, and how to create agents and tools. diff --git a/akd/__init__.py b/akd/__init__.py index b0bb48d9..810dd9c6 100644 --- a/akd/__init__.py +++ b/akd/__init__.py @@ -17,28 +17,15 @@ try: __version__ = _version("akd") except PackageNotFoundError: - __version__ = "0.1.1" + __version__ = "0.2.0" __author__ = "NASA IMPACT" __email__ = "np0069@uah.edu,mr0051@uah.edu" # Core base classes -from akd._base import ( - AbstractBase, - BaseConfig, - InputSchema, - IOSchema, - OutputSchema, - UnrestrictedAbstractBase, -) +from akd._base import AbstractBase, BaseConfig, InputSchema, IOSchema, OutputSchema # Core structures -from akd.structures import ( - ExtractionSchema, - ResearchData, - SearchResultItem, - SingleEstimation, - ToolSearchResult, -) +from akd.structures import SearchResultItem, ToolSearchResult # Tool system from akd.tools._base import BaseTool, BaseToolConfig @@ -50,7 +37,6 @@ "__email__", # Base classes "AbstractBase", - "UnrestrictedAbstractBase", "BaseConfig", "IOSchema", "InputSchema", @@ -60,8 +46,5 @@ "BaseToolConfig", # Core structures "SearchResultItem", - "ResearchData", - "ExtractionSchema", - "SingleEstimation", "ToolSearchResult", ] diff --git a/akd/_base/__init__.py b/akd/_base/__init__.py index ae2844df..89656edd 100644 --- a/akd/_base/__init__.py +++ b/akd/_base/__init__.py @@ -2,18 +2,16 @@ from ._base import ( AbstractBase, - AbstractBaseMeta, - AsyncRunMixin, BaseConfig, InputSchema, IOSchema, OutputSchema, TextInput, TextOutput, - UnrestrictedAbstractBase, ) +from .config_binding import ConfigBindingMixin from .errors import HumanInputRequired -from .exposure import ParamExposureMixin, exposed_param +from .protocols import AKDExecutable, AKDRunContext, AKDTool, RunContextProtocol from .session import AgentSession, BaseSession, ToolSession from .streaming import ( CompletedEvent, @@ -44,13 +42,12 @@ ToolResultEventData, ) from .structures import HumanResponse, RunContext -from .tool_calling import ToolCall, ToolCallingMixin, ToolResult +from .tool_calling import ToolCall, ToolResult +from .validation import validate_input, validate_output, validate_schema __all__ = [ # Base classes "AbstractBase", - "UnrestrictedAbstractBase", - "AsyncRunMixin", # Schema classes "IOSchema", "InputSchema", @@ -59,11 +56,6 @@ "TextOutput", # Config classes "BaseConfig", - # Metadata and decorators - "exposed_param", - "ParamExposureMixin", - # Metaclass - "AbstractBaseMeta", # Streaming "StreamEvent", "StreamEventType", @@ -96,9 +88,19 @@ # Tool calling "ToolCall", "ToolResult", - "ToolCallingMixin", # Context "RunContext", + "AKDRunContext", + # Protocols + "AKDExecutable", + "AKDTool", + "RunContextProtocol", + # Config binding + "ConfigBindingMixin", + # Validation + "validate_input", + "validate_output", + "validate_schema", # Human interaction "HumanResponse", "HumanInputRequired", diff --git a/akd/_base/_base.py b/akd/_base/_base.py index 2c1fedeb..483cc05b 100644 --- a/akd/_base/_base.py +++ b/akd/_base/_base.py @@ -1,26 +1,29 @@ from __future__ import annotations +import asyncio import inspect import types -from abc import ABC, ABCMeta, abstractmethod -from typing import Any, Type, Union, cast, get_args, get_origin +from abc import ABC, abstractmethod +from collections.abc import AsyncIterator +from typing import Any, Generic, Type, TypeVar, Union, get_args, get_origin from loguru import logger -from pydantic import ( - BaseModel, - Field, - PrivateAttr, - ValidationError, - computed_field, - create_model, +from pydantic import BaseModel, Field, PrivateAttr, computed_field, create_model + +from .config_binding import ConfigBindingMixin +from .errors import HumanInputRequired +from .streaming import ( + CompletedEvent, + CompletedEventData, + FailedEvent, + FailedEventData, + RunningEvent, + StartingEvent, + StreamEvent, + StreamEventType, ) - -from akd.utils import get_model_fields, to_snake_case - -from .errors import HumanInputRequired, SchemaValidationError -from .streaming import StreamingMixin from .structures import RunContext -from .utils import AsyncRunMixin +from .validation import validate_input, validate_output class BaseConfig(BaseModel): @@ -107,6 +110,12 @@ def run_context(self) -> RunContext | None: """Get the run context associated with this output.""" return self._run_context + def is_empty(self) -> bool: + """Structurally empty if all fields are None/empty/empty-of-empty (recursive).""" + from akd.utils import is_empty as _is_empty + + return _is_empty(self) + class TextInput(InputSchema): """Simple text-based input schema for unstructured content. @@ -139,167 +148,19 @@ class TextOutput(OutputSchema): content: str = Field(description="The text content response") -def _make_config_property(field_name: str): - """Create a property that references a config field. - - This factory function creates properties at class definition time that delegate - to self.config.field_name, maintaining reference semantics between agent.x and - agent.config.x. Includes fallback to instance __dict__ for pre-init access. - - Args: - field_name: Name of the config field to create a property for - - Returns: - property: A property descriptor with getter/setter - """ - - def getter(self): - # If config doesn't exist yet, fall back to instance attribute - if not hasattr(self, "config") or self.config is None: - return self.__dict__.get(field_name) - return getattr(self.config, field_name) - - def setter(self, value): - # If config doesn't exist yet, set as instance attribute - if not hasattr(self, "config") or self.config is None: - self.__dict__[field_name] = value - else: - setattr(self.config, field_name, value) - - return property(getter, setter) +InSchema = TypeVar("InSchema", bound=InputSchema) +OutSchema = TypeVar("OutSchema", bound=OutputSchema) -def _make_computed_property(field_name: str): - """Create a read-only property for a computed config field. +class AbstractBase(Generic[InSchema, OutSchema], ConfigBindingMixin, ABC): + """Abstract base class for agents and tools. - Computed fields (decorated with @computed_field) are read-only and - dynamically calculated from other config values. + Generic is declared first in the bases so multi-inheritance with + third-party frameworks using the standard ``(Generic[T], ABC)`` + layout (e.g. pydantic-ai) resolves cleanly. - Args: - field_name: Name of the computed field - - Returns: - property: A read-only property descriptor - """ - - def getter(self): - # If config doesn't exist yet, fall back to instance attribute - if not hasattr(self, "config") or self.config is None: - return self.__dict__.get(field_name) - return getattr(self.config, field_name) - - return property(getter) - - -class AbstractBaseMeta(ABCMeta): - """Metaclass that validates required schema attributes and creates config properties.""" - - @staticmethod - def _create_config_properties(target_class, dct): - """Create properties for config fields at class definition time. - - This creates properties that reference self.config.field_name, maintaining - reference semantics between agent.x and agent.config.x. - - Args: - target_class: The class being created - dct: The class dictionary from __new__ - """ - if not hasattr(target_class, "config_schema") or target_class.config_schema is None: - return - - # Create properties for regular model fields - if hasattr(target_class.config_schema, "model_fields"): - for field_name in target_class.config_schema.model_fields.keys(): - # Skip if explicitly defined in this class's dict - if field_name in dct: - continue - # Skip if already a property (from parent or exposed params) - if isinstance(getattr(target_class, field_name, None), property): - continue - # Create property that references self.config.field_name - setattr(target_class, field_name, _make_config_property(field_name)) - - # Create read-only properties for computed fields - if hasattr(target_class.config_schema, "model_computed_fields"): - for field_name in target_class.config_schema.model_computed_fields.keys(): - if field_name in dct: - continue - if isinstance(getattr(target_class, field_name, None), property): - continue - setattr(target_class, field_name, _make_computed_property(field_name)) - - def __new__(mcs, name, bases, dct): - cls = super().__new__(mcs, name, bases, dct) - - # Create config properties at class definition time for ALL classes - # This must happen before early return to ensure base classes get properties too - AbstractBaseMeta._create_config_properties(cls, dct) - - # Skip schema validation for base classes - if name in [ - "AbstractBase", - "UnrestrictedAbstractBase", - "BaseAgent", - "AKDAgent", - "InstructorBaseAgent", - "LiteLLMInstructorBaseAgent", - "BaseTool", - ]: - return cls - - # Check if this class inherits from AbstractBase - if any(isinstance(base, AbstractBaseMeta) for base in bases): - # Validate input_schema - if "input_schema" not in dct and not any(hasattr(base, "input_schema") for base in bases): - raise TypeError(f"{name} must define 'input_schema' class attribute") - - # Validate output_schema - if "output_schema" not in dct and not any(hasattr(base, "output_schema") for base in bases): - raise TypeError(f"{name} must define 'output_schema' class attribute") - - # Validate schema types if they exist - if hasattr(cls, "input_schema") and cls.input_schema is not None: - if not isinstance(cls.input_schema, type) or not issubclass( - cls.input_schema, - (InputSchema, BaseModel), - ): - raise TypeError( - f"{name}.input_schema must be a subclass of InputSchema", - ) - - if hasattr(cls, "output_schema") and cls.output_schema is not None: - output_schema_decl = cls.output_schema - origin = get_origin(output_schema_decl) - if origin in (types.UnionType, Union): - args = get_args(output_schema_decl) - valid_union = bool(args) and all( - isinstance(arg, type) and issubclass(arg, (OutputSchema, BaseModel)) for arg in args - ) - if not valid_union: - raise TypeError( - f"{name}.output_schema union members must be subclasses of OutputSchema", - ) - elif not isinstance(output_schema_decl, type) or not issubclass( - output_schema_decl, - (OutputSchema, BaseModel), - ): - raise TypeError( - f"{name}.output_schema must be a subclass of OutputSchema or a union of them", - ) - - return cls - - -class AbstractBase[ - InSchema: InputSchema, - OutSchema: OutputSchema, -](StreamingMixin, AsyncRunMixin, ABC, metaclass=AbstractBaseMeta): - """ - Abstract base class for agents and tools that interact with a language model. - This class provides the basic structure for an agent or tool that can handle - asynchronous operations, manage memory, and utilize a language model - for generating responses based on user input. + Includes streaming (astream/_astream) and sync run() directly. + Formerly split across StreamingMixin and AsyncRunMixin. """ input_schema: Type[InSchema] @@ -332,210 +193,91 @@ def __class_getitem__(cls, params): attrs["__qualname__"] = f"{cls.__qualname__}[{in_schema.__name__}, {out_label}]" return type(cls.__name__, (cls,), attrs) - def __init__( - self, - config: BaseConfig | BaseModel | None = None, - debug: bool = False, - **kwargs, - ) -> None: - """ - Initializes the BaseAgent with a language model client and memory. - - Args: - debug (bool): If True, enables debug mode for additional logging. - config (BaseModel, optional): Configuration object containing all parameters - debug (bool): If True, enables debug mode for additional logging. - **kwargs: Additional keyword arguments (merged with config) - """ - config = config or (self.config_schema() if self.config_schema else None) or BaseConfig() - self.config = config - self._kwargs = kwargs - self._post_init() - self.debug = debug or getattr(config, "debug", False) - - def _post_init(self) -> None: - """ - Post-initialization hook to perform any additional setup after - the instance has been initialized. - This can be overridden by subclasses for custom behavior. - - Note: Config properties are created by AbstractBaseMeta at class definition time, - not during instance initialization. - """ - for key, value in self._kwargs.items(): - setattr(self, key, value) - - # Set default name from class name if not provided - if getattr(self, "name", None) is None: - self.name = to_snake_case(self.__class__.__name__) - - self.description = (getattr(self, "description", None) or self.__class__.__doc__ or "").strip() - - # Add input/output schema info to description if io_hints is True - if getattr(self, "io_hints", True): - _in_schema = self._input_schema_info - if _in_schema: - self.description += f"\n\nINPUT FIELD DESCRIPTIONS:\n{_in_schema}" - _out_schema = self._output_schema_info - if _out_schema: - self.description += f"\n\nOUTPUT FIELD DESCRIPTIONS:\n{_out_schema}" - - @property - def _input_schema_info(self) -> str: - """ - Extract field names and descriptions from input schema. + # ── Streaming (folded from StreamingMixin) ────────────────────── - Returns: - str: Formatted string with field information, empty if no input schema. + async def astream( + self, + params: Any, + run_context: RunContext | None = None, + **kwargs: Any, + ) -> AsyncIterator[StreamEvent]: + """Public streaming API with input/output validation. + + Wraps _astream() with validation: + - Input validation before any events are yielded + - Output validation on COMPLETED event before yielding """ - # avoid circular dependency - if not hasattr(self, "input_schema") or not self.input_schema: - return "" - - fields = get_model_fields(self.input_schema, skip_no_description=False) - if not fields: - return "" + params = self._validate_input(params) - return "\n".join( - [f"- **{field['name']}**: {field.get('description', field['name'].replace('_', ' '))}" for field in fields], - ) + async for event in self._astream(params, run_context, **kwargs): + if event.event_type == StreamEventType.COMPLETED: + output = event.output + if output is not None: + output = self._validate_output(output) + yield CompletedEvent( + source=event.source, + message=event.message, + data=CompletedEventData(output=output), + run_context=event.run_context, + ) + continue + yield event - @property - def _output_schema_info(self) -> str: - """ - Extract field names and descriptions from output schema. + async def _astream( + self, + params: Any, + run_context: RunContext, + **kwargs: Any, + ) -> AsyncIterator[StreamEvent]: + """Internal streaming implementation. Override for custom streaming. - Returns: - str: Formatted string with field information, empty if no output schema. + Default yields STARTING/RUNNING, calls _arun(), yields COMPLETED/FAILED. """ + class_name = self.__class__.__name__ - # avoid circular dependency - if not hasattr(self, "output_schema") or not self.output_schema: - return "" - schema_decl = self.output_schema - if not isinstance(schema_decl, type): - return "" - fields = get_model_fields(schema_decl, skip_no_description=False) - if not fields: - return "" - - return "\n".join( - [f"- **{field['name']}**: {field.get('description', field['name'].replace('_', ' '))}" for field in fields], + yield StartingEvent( + source=class_name, + message=f"Starting {class_name}", + run_context=run_context, ) - @classmethod - def from_dict(cls, config_dict: dict[str, Any]) -> AbstractBase: - """Create instance from dict, with dynamic config model if needed.""" - debug = config_dict.pop("debug", False) - - # Use existing config_schema or create dynamic one - if cls.config_schema is None and config_dict: - fields = {k: (type(v), v) for k, v in config_dict.items() if v is not None} - cls.config_schema = create_model( # type: ignore[call-arg] - f"{cls.__name__}Config", - __base__=BaseConfig, - **fields, - ) - - config = cls.config_schema(**config_dict) if cls.config_schema and config_dict else None - return cls(config=config, debug=debug) - - def _validate_input(self, params: Any) -> InSchema: - """Validate and convert input parameters.""" - if not isinstance(params, self.input_schema): - if isinstance(params, dict): - try: - params = self.input_schema(**params) - except ValidationError as e: - raise SchemaValidationError(f"Invalid input parameters: {e}") from e - else: - raise TypeError( - f"params must be an instance of {self.input_schema.__name__}", - ) - return params - - def _validate_output(self, output: Any) -> OutSchema: - """Validate output against schema.""" - schema_decl = self.output_schema - origin = get_origin(schema_decl) - if origin in (types.UnionType, Union): - args = [arg for arg in get_args(schema_decl) if isinstance(arg, type) and issubclass(arg, BaseModel)] - if not any(isinstance(output, arg) for arg in args): - raise TypeError( - "Output must be an instance of one of: " + ", ".join(arg.__name__ for arg in args), - ) - return output - if not isinstance(output, schema_decl): - raise TypeError( - f"Output must be an instance of {schema_decl.__name__}", + try: + yield RunningEvent( + source=class_name, + message=f"Running {class_name}", + run_context=run_context, ) - return output - async def arun( - self, - params: InSchema, - **kwargs, - ) -> OutSchema: - """ - Runs the agent with the provided parameters asynchronously. - Args: - params (InSchema): The structured input parameters for the agent. - **kwargs: Additional keyword arguments. - Returns: - OutSchema: The output from the agent after processing the input. - """ + output = await self._arun(params, **kwargs) - params = self._validate_input(params) - if self.debug: - logger.debug( - f"Running {self.__class__.__name__} with params: {params}", + yield CompletedEvent( + source=class_name, + message=f"Completed {class_name}", + data=CompletedEventData(output=output), + run_context=run_context, ) - output = None - try: - output = await self._arun(params, **kwargs) - output = self._validate_output(output) - except HumanInputRequired: - logger.warning(f"{self.__class__.__name__}: HumanInputRequired (flow control)") - raise except Exception as e: - logger.error(f"Error running {self.__class__.__name__}: {e}") + yield FailedEvent( + source=class_name, + message=f"Failed: {e!s}", + data=FailedEventData(error=str(e), error_type=type(e).__name__), + run_context=run_context, + ) raise - return output - - @abstractmethod - async def _arun( - self, - params: InSchema, - **kwargs, - ) -> OutSchema: - """Internal method to run the agent with the provided parameters asynchronously. - Args: - params (InSchema): The structured input parameters for the agent. - **kwargs: Additional keyword arguments. - Returns: - OutSchema: The output from the agent after processing the input. - """ - raise NotImplementedError() + # ── Sync wrapper (folded from AsyncRunMixin) ───────────────────── -class UnrestrictedAbstractBase[ - InSchema: BaseModel, - OutSchema: BaseModel, -](StreamingMixin, AsyncRunMixin, ABC, metaclass=AbstractBaseMeta): - """ - Abstract base class for agents and tools that interact with a language model. - This class provides the basic structure for an agent or tool that can handle - asynchronous operations, manage memory, and utilize a language model - for generating responses based on user input. - - This class does not enforce input and output schema types, allowing for more flexibility - in the types of parameters and outputs used. - It is intended for use cases where strict type checking is not required. - It is recommended to use this class only when necessary, as it bypasses the type safety - provided by the schema validation in the AbstractBase class. - """ - - config_schema: Type[BaseModel] | None = None + def run(self, *args: Any, **kwargs: Any) -> Any: + """Run arun() synchronously.""" + try: + loop = asyncio.get_running_loop() + if loop.is_running(): + future = asyncio.ensure_future(self.arun(*args, **kwargs)) + return loop.run_until_complete(future) + else: + return asyncio.run(self.arun(*args, **kwargs)) + except RuntimeError: + return asyncio.run(self.arun(*args, **kwargs)) def __init__( self, @@ -543,46 +285,70 @@ def __init__( debug: bool = False, **kwargs, ) -> None: - """ - Initializes the BaseAgent with a language model client and memory. + """Initialize the agent/tool. - Args: - debug (bool): If True, enables debug mode for additional logging. - config (BaseModel, optional): Configuration object containing all parameters - debug (bool): If True, enables debug mode for additional logging. - **kwargs: Additional keyword arguments (merged with config) + Validates required schemas at instantiation time. Abstract and + intermediate base classes (e.g. BaseAgent, AKDAgent) can be defined + without concrete schemas, but can only be instantiated through + concrete subclasses that do define them. """ - debug = getattr(config, "debug", False) or debug - self.debug = debug + self._validate_schemas() + config = config or (self.config_schema() if self.config_schema else None) or BaseConfig() self.config = config self._kwargs = kwargs self._post_init() + self.debug = debug or getattr(config, "debug", False) - def _post_init(self) -> None: + def _validate_schemas(self) -> None: + """Validate input_schema / output_schema are defined and of proper types. + + Called at __init__ time. Raises TypeError if schemas are missing or + invalid — which naturally prevents instantiation of abstract/intermediate + classes without requiring hardcoded class name checks. """ - Post-initialization hook to perform any additional setup after - the instance has been initialized. - This can be overridden by subclasses for custom behavior. + cls_name = type(self).__name__ + + # input_schema + input_schema = getattr(self, "input_schema", None) + if input_schema is None: + raise TypeError(f"{cls_name} must define 'input_schema' class attribute") + if not isinstance(input_schema, type) or not issubclass(input_schema, (InputSchema, BaseModel)): + raise TypeError(f"{cls_name}.input_schema must be a subclass of InputSchema") + + # output_schema (supports unions) + output_schema = getattr(self, "output_schema", None) + if output_schema is None: + raise TypeError(f"{cls_name} must define 'output_schema' class attribute") + origin = get_origin(output_schema) + if origin in (types.UnionType, Union): + args = get_args(output_schema) + if not args or not all( + isinstance(arg, type) and issubclass(arg, (OutputSchema, BaseModel)) for arg in args + ): + raise TypeError(f"{cls_name}.output_schema union members must be subclasses of OutputSchema") + elif not isinstance(output_schema, type) or not issubclass(output_schema, (OutputSchema, BaseModel)): + raise TypeError(f"{cls_name}.output_schema must be a subclass of OutputSchema or a union of them") + + def _post_init(self) -> None: + """Post-initialization hook. Subclasses override for custom behavior. - Note: Config properties are created by AbstractBaseMeta at class definition time, - not during instance initialization. + Config properties are created by ConfigBindingMixin at class definition time. + Name/description/IO hints are bound via ConfigBindingMixin._bind_metadata(). """ for key, value in self._kwargs.items(): setattr(self, key, value) - # Set default name from class name if not provided - if getattr(self, "name", None) is None: - self.name = to_snake_case(self.__class__.__name__) + self._bind_metadata() @classmethod - def from_dict(cls, config_dict: dict[str, Any]) -> UnrestrictedAbstractBase: + def from_dict(cls, config_dict: dict[str, Any]) -> AbstractBase: """Create instance from dict, with dynamic config model if needed.""" debug = config_dict.pop("debug", False) # Use existing config_schema or create dynamic one if cls.config_schema is None and config_dict: fields = {k: (type(v), v) for k, v in config_dict.items() if v is not None} - cls.config_schema = create_model( + cls.config_schema = create_model( # type: ignore[call-arg] f"{cls.__name__}Config", __base__=BaseConfig, **fields, @@ -593,15 +359,11 @@ def from_dict(cls, config_dict: dict[str, Any]) -> UnrestrictedAbstractBase: def _validate_input(self, params: Any) -> InSchema: """Validate and convert input parameters.""" - if not isinstance(params, BaseModel): - raise TypeError("params must be an instance of pydantic BaseModel") - return cast(InSchema, params) + return validate_input(self.input_schema, params) def _validate_output(self, output: Any) -> OutSchema: - """Validate and convert input parameters.""" - if not isinstance(output, BaseModel): - raise TypeError("output must be an instance of pydantic BaseModel") - return cast(OutSchema, output) + """Validate output against schema.""" + return validate_output(self.output_schema, output) async def arun( self, @@ -656,5 +418,4 @@ async def _arun( "IOSchema", "InputSchema", "OutputSchema", - "UnrestrictedAbstractBase", ] diff --git a/akd/_base/config_binding.py b/akd/_base/config_binding.py new file mode 100644 index 00000000..17eba902 --- /dev/null +++ b/akd/_base/config_binding.py @@ -0,0 +1,163 @@ +"""Opt-in config and metadata binding mixin. + +Provides config property binding (``agent.x`` → ``agent.config.x``) and +metadata binding (name, description, IO hints) without requiring +``AbstractBaseMeta`` or inheriting from ``AbstractBase``. + +Framework adapters that inherit a third-party agent class directly +(e.g. ``class PydanticAIAgent(pydantic_ai.Agent)``) can mix this in +to get akd-style config access and description building. + +Config properties are collision-aware: fields already defined on the class +(e.g. set by the framework's ``__init__``) are skipped to avoid shadowing. +""" + +from __future__ import annotations + +import types +from typing import Any, Union, get_args, get_origin + +from pydantic import BaseModel + +from akd.utils import get_model_fields, to_snake_case + + +def _make_config_property(field_name: str) -> property: + """Create a property that delegates to ``self.config.``.""" + + def getter(self: Any) -> Any: + if not hasattr(self, "config") or self.config is None: + return self.__dict__.get(field_name) + return getattr(self.config, field_name) + + def setter(self: Any, value: Any) -> None: + if not hasattr(self, "config") or self.config is None: + self.__dict__[field_name] = value + else: + setattr(self.config, field_name, value) + + return property(getter, setter) + + +def _make_computed_property(field_name: str) -> property: + """Create a read-only property for a computed config field.""" + + def getter(self: Any) -> Any: + if not hasattr(self, "config") or self.config is None: + return self.__dict__.get(field_name) + return getattr(self.config, field_name) + + return property(getter) + + +def _format_single_schema(schema: type) -> str: + """Format one schema's fields as bullet lines.""" + fields = get_model_fields(schema, skip_no_description=False) + if not fields: + return "" + return "\n".join( + f"- **{field['name']}**: {field.get('description', field['name'].replace('_', ' '))}" for field in fields + ) + + +def _format_schema_fields(schema: Any) -> str: + """Format schema field names + descriptions as a string. + + Supports single types and union types (e.g. ``A | B``). For unions, + each branch is labeled with its docstring and fields are listed under it. + """ + if schema is None: + return "" + + origin = get_origin(schema) + if origin in (types.UnionType, Union): + branches = [arg for arg in get_args(schema) if isinstance(arg, type) and issubclass(arg, BaseModel)] + if not branches: + return "" + parts = [] + for branch in branches: + doc = (branch.__doc__ or branch.__name__).strip().split("\n")[0] + lines = _format_single_schema(branch) + if lines: + indented = "\n".join(f" {line}" for line in lines.splitlines()) + parts.append(f"**{branch.__name__}**: {doc}\n{indented}") + else: + parts.append(f"**{branch.__name__}**: {doc}") + return "\n".join(parts) + + if isinstance(schema, type): + return _format_single_schema(schema) + + return "" + + +class ConfigBindingMixin: + """Opt-in config property binding + metadata (name, description, IO hints). + + Config properties are created at class definition time via + ``__init_subclass__``. Properties that conflict with existing class + attributes (e.g. framework-owned fields like ``retries``, ``name``) + are skipped automatically. + + Call ``_bind_metadata()`` in your ``__init__`` after ``self.config`` + is set to populate ``self.name`` and ``self.description``. + + Example:: + + class PydanticAIAgent(ConfigBindingMixin, pydantic_ai.Agent): + config_schema = MyConfig + + def __init__(self, config=None, **kw): + self.config = config or self.config_schema() + self._bind_metadata() + pydantic_ai.Agent.__init__(self, system_prompt=self.description, ...) + """ + + def __init_subclass__(cls, **kwargs: Any) -> None: + super().__init_subclass__(**kwargs) + config_cls = getattr(cls, "config_schema", None) + if config_cls is None or not isinstance(config_cls, type) or not issubclass(config_cls, BaseModel): + return + + # Regular model fields + for field_name in getattr(config_cls, "model_fields", {}): + if field_name in cls.__dict__: + continue + existing = getattr(cls, field_name, None) + if existing is not None and not isinstance(existing, property): + continue + setattr(cls, field_name, _make_config_property(field_name)) + + # Computed fields (read-only) + for field_name in getattr(config_cls, "model_computed_fields", {}): + if field_name in cls.__dict__: + continue + existing = getattr(cls, field_name, None) + if existing is not None and not isinstance(existing, property): + continue + setattr(cls, field_name, _make_computed_property(field_name)) + + def _bind_metadata(self) -> None: + """Populate ``self.name`` and ``self.description`` from config + class info. + + Call this in ``__init__`` after ``self.config`` is set. + """ + if not getattr(self, "name", None): + self.name = to_snake_case(type(self).__name__) + + config = getattr(self, "config", None) + desc = (getattr(config, "description", None) or (type(self).__doc__ or "")).strip() + + io_hints = getattr(config, "io_hints", True) if config else True + if io_hints: + in_info = _format_schema_fields(getattr(self, "input_schema", None)) + out_info = _format_schema_fields(getattr(self, "output_schema", None)) + if in_info: + desc += f"\n\nINPUT FIELD DESCRIPTIONS:\n{in_info}" + if out_info: + desc += f"\n\nOUTPUT FIELD DESCRIPTIONS:\n{out_info}" + + self.description = desc + + +__all__ = ["ConfigBindingMixin"] diff --git a/akd/_base/exposure.py b/akd/_base/exposure.py deleted file mode 100644 index bf172798..00000000 --- a/akd/_base/exposure.py +++ /dev/null @@ -1,257 +0,0 @@ -import types -from typing import Any, Callable, get_args, get_type_hints, overload - -from loguru import logger -from pydantic import BaseModel, Field - - -def get_type_from_property(prop: property) -> tuple[Any | None, str]: - """Extract type annotation from property. - - Tries getter return type first, then setter parameter type. - - Args: - prop: Property to extract type from - - Returns: - Tuple of (type_hint, source) where: - - type_hint: The actual type annotation object (e.g., int, str, Union[int, str]) - - source: 'fget' (from getter), 'fset' (from setter), or 'none' (not found) - """ - # Try getter return type first - try: - hints = get_type_hints(prop.fget) - if "return" in hints: - return hints["return"], "fget" - except Exception as e: - logger.warning(f"Failed to get type hints from fget (getter) for {prop}: {e}") - - # Try setter parameter type - if prop.fset: - try: - hints = get_type_hints(prop.fset) - params = [k for k in hints.keys() if k != "return"] - if params: - return hints[params[0]], "fset" - except Exception as e: - logger.warning(f"Failed to get type hints from fset (setter) for {prop}: {e}") - - return None, "none" - - -class ExposedParam(BaseModel): - """Metadata for exposed parameters.""" - - name: str = Field( - description="Parameter name", - ) - description: str = Field( - default="", - description="Human readable description", - ) - extra: dict[str, Any] | None = Field( - default_factory=dict, - description="Additional metadata for the parameter", - ) - - -class ExposedParamRuntimeInfo(ExposedParam): - """Complete runtime information about an exposed parameter. - - Inherits name, description, and extra from ExposedParam. - Adds runtime-specific fields like type, editability, and current value. - """ - - type_: str = Field( - description="Parameter type as string (e.g., 'str', 'int', 'float')", - ) - type_source: str = Field( - description="Source of type inference: 'fget' (getter annotation), 'fset' (setter annotation), 'runtime' (inferred from value), or 'none' (unknown)", - ) - editable: bool = Field( - description="Whether the parameter has a setter and can be modified", - ) - current_value: Any | None = Field( - default=None, - description="Current runtime value of the parameter (only populated when include_values=True)", - ) - - -class ValidatedProperty(property): - """Property subclass that automatically validates setter arguments against type hints. - - When a setter is added via the .setter() method, it's automatically wrapped - with type validation logic that checks the value against type annotations. - """ - - def setter(self, fset): - """Override setter to add automatic type validation.""" - original_fset = fset - - def validated_setter(instance, value): - # Get expected type from property getter first - expected_type, _ = get_type_from_property(self) - - # If no type from getter, try the new setter being added - if expected_type is None: - try: - hints = get_type_hints(original_fset) - params = [k for k in hints.keys() if k != "return"] - if params: - expected_type = hints[params[0]] - except Exception: - pass - - # Validate type if we found one - if expected_type is not None: - # Get origin type for generics (List[int] -> list) - origin = getattr(expected_type, "__origin__", None) - - # Check if it's a Union type (typing.Union or Python 3.10+ int | str) - is_union = origin is not None or isinstance(expected_type, types.UnionType) - - if is_union: - # Handle Union types (e.g., Union[int, str] or int | str) - type_args = get_args(expected_type) - if type_args and not isinstance(value, type_args): - param_name = getattr(self.fget, "_exposed_meta", None) - name = param_name.name if param_name else "parameter" - type_names = ", ".join(t.__name__ if hasattr(t, "__name__") else str(t) for t in type_args) - raise TypeError( - f"Parameter '{name}' expects one of [{type_names}], got {type(value).__name__}", - ) - else: - # Simple type check - if not isinstance(value, expected_type): - param_name = getattr(self.fget, "_exposed_meta", None) - name = param_name.name if param_name else "parameter" - type_name = getattr(expected_type, "__name__", str(expected_type)) - raise TypeError( - f"Parameter '{name}' expects {type_name}, got {type(value).__name__}", - ) - - # Call original setter - return original_fset(instance, value) - - return super().setter(validated_setter) - - -@overload -def exposed_param[F: Callable](_func: F, /) -> property: ... - - -@overload -def exposed_param[F: Callable]( - *, - description: str = "", - **kwargs, -) -> Callable[[F], property]: ... - - -def exposed_param[F: Callable]( - _func: F | None = None, - /, - *, - description: str = "", - **kwargs, -) -> property | Callable[[F], property]: - """ - Decorator to mark a property as exposed to external systems (UI, backend, API, etc.). - - This decorator combines @property functionality with parameter exposure metadata, - making the parameter accessible and configurable from external systems. - - Can be used with or without parentheses: - @exposed_param - @exposed_param(description="...", extra_field="...") - - Args: - _func: The function being decorated (positional-only, internal use) - description: Human readable description of the parameter - **kwargs: Additional metadata stored in 'extra' field (e.g., constraints, hints) - """ - - def decorator(func: F) -> property: - # Priority: explicit description > function docstring > prettified function name - final_description = description or (func.__doc__ or "").strip() or func.__name__.replace("_", " ").title() - - func._exposed_meta = ExposedParam( # type: ignore[attr-defined] - name=func.__name__, - description=final_description, - extra=kwargs or {}, - ) - return ValidatedProperty(func) # type: ignore[return-value] - - if _func is not None: - return decorator(_func) - - return decorator - - -class ParamExposureMixin: - """Mixin to add parameter exposure capabilities to agents.""" - - def get_exposed_params(self, include_values: bool = False) -> list[ExposedParamRuntimeInfo]: - """ - Get metadata about exposed parameters. - - Args: - include_values: If True, include current runtime values in the output - - Returns: - List of ExposedParamRuntimeInfo objects, each containing: - - name: Parameter name - - description: Human readable description - - type_: Python type name (from type hints or runtime value) - - type_source: Source of type inference - - editable: Whether the parameter has a setter - - extra: Additional metadata from decorator - - current_value: Current value (if include_values=True) - """ - exposed = [] - for attr_name in dir(self): - attr = getattr(type(self), attr_name, None) - if not (isinstance(attr, property) and hasattr(attr.fget, "_exposed_meta")): - continue - - meta = attr.fget._exposed_meta - - # Get type from hints using shared utility - type_hint, type_source = get_type_from_property(attr) - current_value = None - - # Convert type hint to string name - if type_hint is not None: - param_type = getattr(type_hint, "__name__", str(type_hint)) - else: - param_type = "unknown" - - # Fallback to runtime value if no type hint - if param_type == "unknown": - try: - current_value = getattr(self, attr_name) - param_type = type(current_value).__name__ - type_source = "runtime" - except Exception: - param_type = "unknown" - - # Get current value if requested and not already fetched - if include_values and current_value is None: - try: - current_value = getattr(self, attr_name) - except Exception: - current_value = None - - exposed.append( - ExposedParamRuntimeInfo( - name=attr_name, - description=meta.description, - type_=param_type, - type_source=type_source, - editable=attr.fset is not None, - extra=meta.extra, - current_value=current_value if include_values else None, - ), - ) - - return exposed diff --git a/akd/_base/protocols.py b/akd/_base/protocols.py new file mode 100644 index 00000000..f67f7504 --- /dev/null +++ b/akd/_base/protocols.py @@ -0,0 +1,100 @@ +"""Type protocols for the AKD framework. + +Structural supertypes defining the contracts for AKD executables +(agents and tools). Concrete classes like ``AbstractBase``, ``BaseAgent``, +and ``BaseTool`` satisfy these protocols structurally — no explicit +Protocol inheritance required. + +Framework adapters can satisfy these protocols either by inheriting +the concrete implementation kit (``BaseAgent``, ``BaseTool``) or by +implementing the methods directly on a framework-native class +(e.g. ``class PydanticAIAgent(pydantic_ai.Agent)``). +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable +from typing import TYPE_CHECKING, Any, Literal, Protocol, runtime_checkable + +if TYPE_CHECKING: + from .streaming import StreamEvent + from .structures import RunContext + + +@runtime_checkable +class TokenCounts(Protocol): + """Structural shape for usage/token-counter objects. + + Both ``akd.RunUsage`` (BaseModel) and ``pydantic_ai.RunUsage`` + (dataclass) satisfy this structurally. + """ + + input_tokens: int + output_tokens: int + requests: int + + +@runtime_checkable +class RunContextProtocol(Protocol): + """Structural supertype for run-context-like objects. + + Both ``akd.RunContext`` and ``pydantic_ai.RunContext`` satisfy this + structurally for the overlapping fields below. For akd-specific + fields (``human_response``, ``extra``) or framework-specific fields + (``deps``, ``tool_call_id``, etc.), narrow to the concrete type. + """ + + messages: list[Any] | None + usage: TokenCounts + run_id: str | None + + +@runtime_checkable +class AKDExecutable(Protocol): + """Contract shared by all AKD runnables — agents and tools alike. + + Captures: typed I/O schemas, typed config, name/description metadata, + and the async ``arun``/``astream`` entry points. + """ + + input_schema: type + output_schema: type + config_schema: type + name: str | None + description: str | None + + async def arun( + self, + params: Any, + run_context: RunContext | None = None, + **kwargs: Any, + ) -> Any: ... + + async def astream( + self, + params: Any, + run_context: RunContext | None = None, + **kwargs: Any, + ) -> AsyncIterator[StreamEvent]: ... + + +@runtime_checkable +class AKDTool(AKDExecutable, Protocol): + """Tool-specific extension — adds function export.""" + + def as_function( + self, + mode: Literal["python", "json"] | None = None, + ) -> Callable[..., Awaitable[Any]]: ... + + +# Disambiguating alias — use when pydantic_ai.RunContext is also in scope. +from .structures import RunContext as AKDRunContext # noqa: E402 + +__all__ = [ + "AKDExecutable", + "AKDRunContext", + "AKDTool", + "RunContextProtocol", + "TokenCounts", +] diff --git a/akd/_base/structures.py b/akd/_base/structures.py index be5f3022..16d2d244 100644 --- a/akd/_base/structures.py +++ b/akd/_base/structures.py @@ -103,9 +103,13 @@ class RunContext(BaseModel): default=None, description="Human response for resumption after HUMAN_INPUT_REQUIRED", ) - messages: list[dict[str, Any]] | None = Field( + messages: list[Any] | None = Field( default=None, - description="Conversation history for resumption", + description=( + "Conversation history for resumption. AKD's own BaseAgent populates " + "these as OpenAI-style {role, content} dicts; adapters (pydantic-ai, " + "langchain, etc.) may carry their native message types instead." + ), ) run_id: str | None = Field( default=None, diff --git a/akd/_base/tool_calling.py b/akd/_base/tool_calling.py index eb701220..4d92b8e2 100644 --- a/akd/_base/tool_calling.py +++ b/akd/_base/tool_calling.py @@ -1,17 +1,16 @@ -"""Tool calling support for akd agents.""" +"""Tool call / result data models for akd agents. + +Provider-specific adapters convert their native tool call formats +into these canonical types. +""" from __future__ import annotations -import asyncio from datetime import UTC, datetime -from typing import TYPE_CHECKING, Any +from typing import Any -from loguru import logger from pydantic import BaseModel, Field -if TYPE_CHECKING: - from akd.tools._base import BaseTool - class ToolCall(BaseModel): """Tool call request (normalized from any provider). @@ -40,106 +39,4 @@ class ToolResult(BaseModel): error: str | None = Field(default=None, description="Error message if failed") -class ToolCallingMixin: - """Mixin providing reusable tool execution helpers. - - Requires the inheriting class to have: - - self.tools: list[BaseTool] - - Example: - class MyAgent(ToolCallingMixin, BaseAgent): - async def _astream(self, params, context, **kwargs): - # Parse provider format → ToolCall - tool_call = ToolCall( - tool_call_id=tc.id, - tool_name=tc.function.name, - arguments=json.loads(tc.function.arguments), - ) - result = await self._execute_tool(tool_call) - """ - - def _find_tool( - self, - name: str, - tools: list[BaseTool] | None = None, - ) -> BaseTool | None: - """Find a tool by name. - - Args: - name: The name of the tool to find (checks both tool.name and class name) - tools: Optional list of tools to search. Defaults to self.tools. - - Returns: - The matching tool instance, or None if not found - """ - tools = tools or self.tools - return next( - (t for t in tools if t.name == name or t.__class__.__name__ == name), - None, - ) - - async def _execute_tool( - self, - tool_call: ToolCall, - tools: list[BaseTool] | None = None, - ) -> ToolResult: - """Execute a single tool call. - - Args: - tool_call: Normalized tool call request - tools: Optional list of tools to search. Defaults to self.tools. - - Returns: - ToolResult with content or error - """ - tool = self._find_tool(tool_call.tool_name, tools=tools) - if not tool: - return ToolResult( - tool_call_id=tool_call.tool_call_id, - tool_name=tool_call.tool_name, - content=None, - error=f"Unknown tool: {tool_call.tool_name}", - ) - - try: - input_obj = tool.input_schema(**tool_call.arguments) - result = await tool.arun(input_obj) - # Use mode='json' to ensure JSON-serializable types (HttpUrl → str, datetime → ISO string) - content = result.model_dump(mode="json") if hasattr(result, "model_dump") else result - return ToolResult( - tool_call_id=tool_call.tool_call_id, - tool_name=tool_call.tool_name, - content=content, - ) - except Exception as e: - logger.exception(f"Tool '{tool_call.tool_name}' failed with args {tool_call.arguments}") - return ToolResult( - tool_call_id=tool_call.tool_call_id, - tool_name=tool_call.tool_name, - content=None, - error=str(e), - ) - - async def _execute_tools_parallel( - self, - tool_calls: list[ToolCall], - tools: list[BaseTool] | None = None, - ) -> list[ToolResult]: - """Execute multiple tool calls in parallel. - - When an LLM returns multiple tool calls in a single response, - they are independent by definition and can be executed concurrently. - - Args: - tool_calls: List of normalized tool calls - tools: Optional list of tools to search. Defaults to self.tools. - - Returns: - List of ToolResults in same order as input - """ - return list( - await asyncio.gather(*[self._execute_tool(tc, tools=tools) for tc in tool_calls]), - ) - - -__all__ = ["ToolCall", "ToolResult", "ToolCallingMixin"] +__all__ = ["ToolCall", "ToolResult"] diff --git a/akd/_base/validation.py b/akd/_base/validation.py new file mode 100644 index 00000000..2cafa38e --- /dev/null +++ b/akd/_base/validation.py @@ -0,0 +1,56 @@ +"""Standalone input/output validation functions. + +Used by both AbstractBase (internally) and framework adapters that +don't inherit AbstractBase (e.g. PydanticAIAgent). +""" + +from __future__ import annotations + +import types +from typing import Any, Union, get_args, get_origin + +from pydantic import BaseModel, ValidationError + +from .errors import SchemaValidationError + + +def validate_schema(schema: type, value: Any, *, coerce_dict: bool = False) -> Any: + """Validate value against a schema. Supports union types (A | B). + + If coerce_dict=True, attempts schema(**value) when value is a dict. + """ + origin = get_origin(schema) + if origin in (types.UnionType, Union): + args = [arg for arg in get_args(schema) if isinstance(arg, type) and issubclass(arg, BaseModel)] + if isinstance(value, tuple(args)): + return value + if coerce_dict and isinstance(value, dict): + for arg in args: + try: + return arg(**value) + except Exception: + continue + raise SchemaValidationError(f"Dict doesn't match any of: {[a.__name__ for a in args]}") + raise TypeError("Must be an instance of one of: " + ", ".join(a.__name__ for a in args)) + + if isinstance(value, schema): + return value + if coerce_dict and isinstance(value, dict): + try: + return schema(**value) + except ValidationError as e: + raise SchemaValidationError(f"Invalid parameters: {e}") from e + raise TypeError(f"Must be an instance of {schema.__name__}") + + +def validate_input(input_schema: type, params: Any) -> Any: + """Validate input params — coerces dicts, supports unions.""" + return validate_schema(input_schema, params, coerce_dict=True) + + +def validate_output(output_schema: type, output: Any) -> Any: + """Validate output — supports unions, no dict coercion.""" + return validate_schema(output_schema, output) + + +__all__ = ["validate_input", "validate_output", "validate_schema"] diff --git a/akd/agents/_base/__init__.py b/akd/agents/_base/__init__.py index c2915998..a4c83b35 100644 --- a/akd/agents/_base/__init__.py +++ b/akd/agents/_base/__init__.py @@ -6,6 +6,7 @@ from ._base import ( Agent, AKDAgent, + AKDBaseAgent, BaseAgent, BaseAgentConfig, InstructorBaseAgent, @@ -15,6 +16,7 @@ __all__ = [ "AKDAgent", + "AKDBaseAgent", "Agent", "BaseAgent", "BaseAgentConfig", diff --git a/akd/agents/_base/_base.py b/akd/agents/_base/_base.py index 3b0163e7..1df51790 100644 --- a/akd/agents/_base/_base.py +++ b/akd/agents/_base/_base.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import copy import json import uuid @@ -19,13 +20,12 @@ BaseConfig, InputSchema, OutputSchema, - ParamExposureMixin, RunContext, StreamEvent, StreamEventType, TextOutput, ToolCall, - ToolCallingMixin, + ToolResult, ) from akd._base.errors import ( HumanInputRequired, @@ -59,7 +59,6 @@ from akd.configs.prompts import DEFAULT_SYSTEM_PROMPT from akd.tools._base import BaseTool from akd.tools.human import HumanTool, HumanToolInput -from akd.utils import get_model_fields from .output_routing import OutputRoutingMixin @@ -207,7 +206,7 @@ def validate_reasoning_params(self): class BaseAgent[ InSchema: InputSchema, OutSchema: OutputSchema, -](AbstractBase, ParamExposureMixin): +](AbstractBase): """Framework-agnostic base class for chat agents. Provides session management, validation, lifecycle, and streaming template. @@ -242,37 +241,27 @@ def output_schema_resolved(self) -> list[type[OutputSchema]]: return unique @property - def _output_schema_info(self) -> str: - """Override to describe each branch of union output schemas.""" - base = super()._output_schema_info - schemas = self.output_schema_resolved - if len(schemas) <= 1: - return base - - parts = [] - for schema in schemas: - doc = (schema.__doc__ or schema.__name__).strip().split("\n")[0] - fields = get_model_fields(schema, skip_no_description=False) - field_lines = "\n".join( - f" - **{f['name']}**: {f.get('description', f['name'].replace('_', ' '))}" for f in fields - ) - parts.append(f"**{schema.__name__}**: {doc}\n{field_lines}") - union_info = "\n".join(parts) - return f"{base}\n{union_info}" if base else union_info + def effective_system_prompt(self) -> str: + """System prompt actually sent to the LLM — base prompt + agent description. - @property - def _system_prompt(self) -> str: - """Enhanced system prompt with agent description.""" + Computed on every access from ``self.system_prompt`` + ``self.description``. + The underlying ``config.system_prompt`` is never mutated. + """ content = self.system_prompt if self.description: content += f"\n\nAGENT DESCRIPTION:\n{self.description}" return content + @property + def _system_prompt(self) -> str: + """Backward-compat alias for ``effective_system_prompt``.""" + return self.effective_system_prompt + def _default_system_message(self) -> dict[str, str]: """Return default system message.""" return { "role": "system", - "content": self._system_prompt, + "content": self.effective_system_prompt, } def _build_run_context(self, run_context: RunContext | None) -> RunContext: @@ -456,7 +445,7 @@ async def _astream( class AKDAgent[ InSchema: InputSchema, OutSchema: OutputSchema, -](ToolCallingMixin, OutputRoutingMixin, BaseAgent): +](OutputRoutingMixin, BaseAgent): """Built-in agent using LiteLLM + instructor. Uses LiteLLM for completion calls and instructor for structured output. @@ -472,6 +461,65 @@ def __init__( super().__init__(config=config, debug=debug) self.client = instructor.from_litellm(acompletion) + # ── Tool execution helpers (folded from ToolCallingMixin) ─────── + + def _find_tool( + self, + name: str, + tools: list[BaseTool] | None = None, + ) -> BaseTool | None: + """Find a tool by name (checks both tool.name and class name).""" + tools = tools or self.tools + return next( + (t for t in tools if t.name == name or t.__class__.__name__ == name), + None, + ) + + async def _execute_tool( + self, + tool_call: ToolCall, + tools: list[BaseTool] | None = None, + ) -> ToolResult: + """Execute a single tool call, returning normalized ToolResult.""" + tool = self._find_tool(tool_call.tool_name, tools=tools) + if not tool: + return ToolResult( + tool_call_id=tool_call.tool_call_id, + tool_name=tool_call.tool_name, + content=None, + error=f"Unknown tool: {tool_call.tool_name}", + ) + + try: + input_obj = tool.input_schema(**tool_call.arguments) + result = await tool.arun(input_obj) + content = result.model_dump(mode="json") if hasattr(result, "model_dump") else result + return ToolResult( + tool_call_id=tool_call.tool_call_id, + tool_name=tool_call.tool_name, + content=content, + ) + except Exception as e: + logger.exception(f"Tool '{tool_call.tool_name}' failed with args {tool_call.arguments}") + return ToolResult( + tool_call_id=tool_call.tool_call_id, + tool_name=tool_call.tool_name, + content=None, + error=str(e), + ) + + async def _execute_tools_parallel( + self, + tool_calls: list[ToolCall], + tools: list[BaseTool] | None = None, + ) -> list[ToolResult]: + """Execute multiple tool calls concurrently.""" + return list( + await asyncio.gather(*[self._execute_tool(tc, tools=tools) for tc in tool_calls]), + ) + + # ── End folded helpers ───────────────────────────────────────── + def _post_init(self): super()._post_init() @@ -933,6 +981,10 @@ async def _run_engine_stream( raise UnexpectedModelBehavior("Non-tool streaming ended without completion") +# Aliases for naming symmetry with framework adapters +# (e.g. OpenAIBaseAgent, PydanticAIBaseAgent — AKDBaseAgent is the akd-native abstract agent) +AKDBaseAgent = BaseAgent + # Backward compatibility aliases LiteLLMInstructorBaseAgent = AKDAgent InstructorBaseAgent = AKDAgent diff --git a/akd/agents/extraction.py b/akd/agents/extraction.py deleted file mode 100644 index d32a5b8d..00000000 --- a/akd/agents/extraction.py +++ /dev/null @@ -1,63 +0,0 @@ -from abc import ABC -from typing import Any, List, Union - -from loguru import logger -from pydantic import BaseModel, Field - -from akd._base import AsyncRunMixin -from akd.agents import LiteLLMInstructorBaseAgent -from akd.structures import ExtractionSchema, SingleEstimation - -from .intents import Intent - - -class ExtractionSchemaMapper(ABC, AsyncRunMixin): - def __init__(self, debug: bool = False) -> None: - self.debug = bool(debug) - - def __call__(self, *args, **kwargs) -> Any: - return self.run(*args, **kwargs) - - -class IntentBasedExtractionSchemaMapper(ExtractionSchemaMapper): - """ - If Intent is ESTIMATION, return a type `List[SingleEstimation]`. - - If GENERAL, return base ExtractionSchema - """ - - async def arun( - self, - intent: Intent, - **kwargs, - ) -> Union[ExtractionSchema, List[SingleEstimation]]: - res = ExtractionSchema - if intent == Intent.ESTIMATION: - res = List[SingleEstimation] - if self.debug: - logger.debug(f"Intent={intent} | Schema={res}") - return res - - -class ExtractionInputSchema(BaseModel): - """Information Extraction input schema""" - - query: str = Field(..., description="Query that is used for answering/extraction") - content: str = Field( - ..., - description="Actual text/content to extract information from", - ) - - -class EstimationExtractionOutputSchema(BaseModel): - """Estimation Extraction output schema""" - - estimations: List[SingleEstimation] = Field( - ..., - description="List of estimations extracted from the query and the content", - ) - - -class EstimationExtractionAgent(LiteLLMInstructorBaseAgent): - input_schema = ExtractionInputSchema - output_schema = EstimationExtractionOutputSchema diff --git a/akd/agents/factory.py b/akd/agents/factory.py deleted file mode 100644 index 409b207b..00000000 --- a/akd/agents/factory.py +++ /dev/null @@ -1,110 +0,0 @@ -from akd.agents import BaseAgentConfig -from akd.agents.extraction import EstimationExtractionAgent -from akd.agents.intents import IntentAgent -from akd.agents.query import FollowUpQueryAgent, QueryAgent -from akd.agents.relevancy import MultiRubricRelevancyAgent, RelevancyAgent -from akd.agents.search import ControlledSearchAgent -from akd.configs.project import CONFIG -from akd.configs.prompts import ( - DEFAULT_SYSTEM_PROMPT, - EXTRACTION_SYSTEM_PROMPT, - INTENT_SYSTEM_PROMPT, - MULTI_RUBRIC_RELEVANCY_SYSTEM_PROMPT, - QUERY_SYSTEM_PROMPT, -) - - -def create_intent_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> IntentAgent: - config = config or BaseAgentConfig( - api_key=CONFIG.model_config_settings.api_keys.openai, - model_name=CONFIG.model_config_settings.model_name, - temperature=CONFIG.model_config_settings.temperature, - system_prompt=INTENT_SYSTEM_PROMPT, - ) - return IntentAgent(config, debug=debug) - - -def create_extraction_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> EstimationExtractionAgent: - config = config or BaseAgentConfig( - api_key=CONFIG.model_config_settings.api_keys.openai, - model_name=CONFIG.model_config_settings.model_name, - temperature=CONFIG.model_config_settings.temperature, - system_prompt=EXTRACTION_SYSTEM_PROMPT, - ) - return EstimationExtractionAgent(config, debug=debug) - - -def create_query_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> QueryAgent: - config = config or BaseAgentConfig( - api_key=CONFIG.model_config_settings.api_keys.openai, - model_name=CONFIG.model_config_settings.model_name, - temperature=CONFIG.model_config_settings.temperature, - system_prompt=QUERY_SYSTEM_PROMPT, - ) - return QueryAgent(config, debug=debug) - - -def create_multi_rubric_relevancy_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> MultiRubricRelevancyAgent: - config = config or BaseAgentConfig( - api_key=CONFIG.model_config_settings.api_keys.openai, - model_name=CONFIG.model_config_settings.model_name, - temperature=CONFIG.model_config_settings.temperature, - system_prompt=MULTI_RUBRIC_RELEVANCY_SYSTEM_PROMPT, - ) - return MultiRubricRelevancyAgent(config, debug=debug) - - -def create_followupquery_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> FollowUpQueryAgent: - """Create a FollowUpQueryAgent instance.""" - config = config or BaseAgentConfig( - api_key=CONFIG.model_config_settings.api_keys.openai, - model_name=CONFIG.model_config_settings.model_name, - temperature=CONFIG.model_config_settings.temperature, - system_prompt=QUERY_SYSTEM_PROMPT, # Reuse query prompt - ) - return FollowUpQueryAgent(config, debug=debug) - - -def create_relevancy_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> RelevancyAgent: - """Create a basic RelevancyAgent instance.""" - config = config or BaseAgentConfig( - api_key=CONFIG.model_config_settings.api_keys.openai, - model_name=CONFIG.model_config_settings.model_name, - temperature=CONFIG.model_config_settings.temperature, - system_prompt=DEFAULT_SYSTEM_PROMPT, # Use default for basic relevancy - ) - return RelevancyAgent(config, debug=debug) - - -def create_lit_agent( - config: BaseAgentConfig | None = None, - debug: bool = False, -) -> ControlledSearchAgent: - """Create a ControlledSearchAgent with default configuration.""" - from akd.agents.search import ControlledSearchAgentConfig - - # Use the new agent's config system - agent_config = ControlledSearchAgentConfig() - if config: - # Map basic config properties to new config if provided - agent_config.debug = debug - - return ControlledSearchAgent(config=agent_config, debug=debug) diff --git a/akd/agents/gap_analysis/__init__.py b/akd/agents/gap_analysis/__init__.py deleted file mode 100644 index a2e43b36..00000000 --- a/akd/agents/gap_analysis/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -from .gap_analysis import GapAgent, GapAgentConfig, GapInputSchema, GapOutputSchema - -__all__ = [ - "GapAgent", - "GapInputSchema", - "GapOutputSchema", - "GapAgentConfig", -] diff --git a/akd/agents/gap_analysis/gap_analysis.py b/akd/agents/gap_analysis/gap_analysis.py deleted file mode 100644 index a4ffe25c..00000000 --- a/akd/agents/gap_analysis/gap_analysis.py +++ /dev/null @@ -1,298 +0,0 @@ -import asyncio -from typing import Dict, List, Tuple - -import networkx as nx -from langchain_openai import ChatOpenAI -from loguru import logger -from networkx.readwrite import json_graph -from pydantic.fields import Field - -from akd._base import InputSchema, OutputSchema -from akd.agents._base import BaseAgent, BaseAgentConfig -from akd.structures import PaperDataItem, SearchResultItem -from akd.tools.scrapers import ( - DoclingScraper, - DoclingScraperConfig, - OmniScraperInputSchema, -) -from akd.tools.search import ( - SemanticScholarSearchTool, - SemanticScholarSearchToolConfig, - SemanticScholarSearchToolInputSchema, -) - -from .graph_utils import add_paper_to_graph, generate_final_answer, select_nodes -from .parsing_utils import ( - create_sections_from_parsed_html, - group_section_titles, - parse_html, -) -from .prompts import GAP_QUERY_MAP -from .structures import ParsedPaper - - -class GapInputSchema(InputSchema): - """Input schema for gap agent""" - - search_results: List[SearchResultItem] = Field( - ..., - description="List of SearchResultItems.", - ) - gap: str = Field(..., description="Type of gap to investigate.") - - -class GapOutputSchema(OutputSchema): - """Output schema for gap agent""" - - __response_field__ = "output" - - model_config = {"arbitrary_types_allowed": True} - output: str = Field(..., description="Final output generated by the agent.") - attributed_source_answers: Dict = Field( - ..., - description="Answers for each node selected from the graph.", - ) - graph: dict | None = Field(default=None, description="Graph created from the ingested papers.") - - -class GapAgentConfig(BaseAgentConfig): - """Configuration for Gap Agent""" - - docling_config: DoclingScraperConfig = Field( - default_factory=lambda: DoclingScraperConfig( - do_table_structure=False, - pdf_mode="fast", - export_type="html", - debug=False, - ), - description="Configuration for Docling Scraper.", - ) - s2_tool_config: SemanticScholarSearchToolConfig = Field( - default_factory=lambda: SemanticScholarSearchToolConfig( - debug=False, - external_id="ARXIV", - fields=[ - "paperId", - "externalIds", - "url", - "title", - "abstract", - "year", - "authors", - "isOpenAccess", - "openAccessPdf", - ], - ), - description="Configuration for S2 Tool.", - ) - - output_graph: bool = Field( - default=False, - description="Whether to output the graph or not.", - ) - - -class GapAgent(BaseAgent): - input_schema = GapInputSchema - output_schema = GapOutputSchema - config_schema = GapAgentConfig - - def _post_init( - self, - ) -> None: - super()._post_init() - self.docling_scraper = DoclingScraper(self.config.docling_config) - self.semantic_search_tool = SemanticScholarSearchTool( - self.config.s2_tool_config, - ) - - self.llm = ChatOpenAI( - model_name=self.config.model_name, - temperature=self.config.temperature, - api_key=self.config.api_key, - ) - - async def _fetch_paper_items( - self, - search_results: List[SearchResultItem], - ) -> Tuple[list[PaperDataItem], list[SearchResultItem]]: - """ - Fetches paper metadata using arXiv IDs extracted from a list of search results. - This function is restricted to Arxiv URLs for now. - - Args: - search_results (List[SearchResultItem]): A list of search result items containing URLs to arXiv papers. - - Returns: - List[PaperDataItem]: A list of paper data items retrieved from the Semantic Scholar tool. - List[SearchResultItem]: The subset of search_results corresponding to the successfully fetched papers. - """ - arxiv_ids = [res.url.path.split("/")[-1].split("v")[0] for res in search_results] - paper_items = await self.semantic_search_tool.fetch_paper_by_external_id( - SemanticScholarSearchToolInputSchema(queries=arxiv_ids), - ) - fetched_paper_ids = [paper_item.external_id for paper_item in paper_items] - skipped_ids = (set(arxiv_ids)) - set(fetched_paper_ids) - search_results = [search_results[i] for i in range(len(search_results)) if arxiv_ids[i] not in skipped_ids] - for paper_item, res in zip(paper_items, search_results): - if paper_item.url is None: - paper_item.url = str(res.url) - return paper_items, search_results - - async def _fetch_parsed_pdfs( - self, - search_results: List[SearchResultItem], - ) -> List[str]: - """ - Parses PDFs from search result items using Docling. - - Args: - search_results (List[SearchResultItem]): A list of search result items containing PDF URLs. - - Returns: - List[str]: A list of parsed PDF contents as strings. - """ - pdf_urls = [res.pdf_url for res in search_results if res is not None] - tasks = [self.docling_scraper.arun(OmniScraperInputSchema(url=url)) for url in pdf_urls] - results = await asyncio.gather(*tasks) - parsed_pdfs = [res.content for res in results] - return parsed_pdfs - - async def _create_parsed_paper( - self, - parsed_pdf: str, - paper_item: PaperDataItem, - ) -> ParsedPaper: - """ - Converts parsed HTML content of a paper and its metadata into a structured ParsedPaper object. - - Args: - parsed_pdf (str): The HTML content extracted and parsed from the paper's PDF. - paper_item (PaperDataItem): PaperDataItem of the parsed_pdf. - - Returns: - ParsedPaper: A structured representation of the paper including section titles, content, and metadata. - """ - parsed_content = parse_html(parsed_pdf) - section_titles = await group_section_titles(parsed_content, self.llm) - sections, clean_section_titles = create_sections_from_parsed_html( - parsed_content, - section_titles, - paper_title=paper_item.title, - ) - parsed_paper = ParsedPaper( - section_titles=clean_section_titles, - sections=sections, - **paper_item.model_dump(), - ) - return parsed_paper - - async def _fetch_parsed_papers( - self, - parsed_pdfs: list[str], - paper_items: List[PaperDataItem], - ) -> List[ParsedPaper]: - """ - Creates structured ParsedPaper objects from parsed PDF content and corresponding paper metadata. - - Args: - parsed_pdfs (list[str]): A list of HTML strings parsed from the original paper PDFs. - paper_items (List[PaperDataItem]): A list of PaperDataItem objects corresponding to each parsed PDF. - - Returns: - List[ParsedPaper]: A list of structured ParsedPaper objects containing organized content and metadata. - """ - tasks = [ - self._create_parsed_paper(parsed_pdf, paper_item) - for parsed_pdf, paper_item in zip(parsed_pdfs, paper_items) - ] - parsed_papers = await asyncio.gather(*tasks) - return parsed_papers - - async def create_graph(self, parsed_papers: List[ParsedPaper]) -> nx.Graph: - """ - Constructs a NetworkX graph from a list of parsed papers by adding their structured content and metadata. - - Args: - parsed_papers (List[ParsedPaper]): A list of ParsedPaper objects containing sectioned content and metadata. - - Returns: - nx.Graph: A graph where nodes and edges represent the structure and relationships within and between papers. - """ - graph = nx.Graph() - for paper in parsed_papers: - graph = await add_paper_to_graph(graph, paper, self.llm) - return graph - - def get_node_data(self, G: nx.Graph, node_id: str) -> Dict: - """ - Retrieves the data associated with a specific node in the graph. - - Args: - G (networkx.Graph): The graph containing nodes with metadata. - node_id (str): The identifier of the node whose data is to be retrieved. - - Returns: - Dict: A dictionary containing all attributes stored for the specified node. - """ - return G.nodes[node_id] - - async def get_response_async( - self, - params: GapInputSchema, - **kwargs, - ) -> GapOutputSchema: - """ - Obtains a response from the language model asynchronously. - - Args: - response_model (Optional[OutputSchema]): - The schema for the response data. If not set, - self.output_schema is used. - - Returns: - OutputSchema: The response from the language model. - """ - search_results = params.search_results - if self.debug: - if params.gap not in GAP_QUERY_MAP.keys(): - logger.debug( - "You are running the gap agent with your own defined gap. Please ensure you have described the gap you want to investigate in detail.", - ) - else: - logger.debug(f"Running gap analysis to investigate {params.gap} gap.") - gap = GAP_QUERY_MAP[params.gap] if params.gap in GAP_QUERY_MAP.keys() else params.gap - paper_items, search_results = await self._fetch_paper_items(search_results) - if self.debug: - logger.debug(f"Fetch {len(paper_items)} papers from semantic scholar.") - parsed_pdfs = await self._fetch_parsed_pdfs(search_results=search_results) - if self.debug: - logger.debug(f"Parsed {len(parsed_pdfs)} pdf documents.") - parsed_papers = await self._fetch_parsed_papers(parsed_pdfs, paper_items) - if self.debug: - logger.debug(f"Fetched {len(parsed_papers)} parsed papers.") - graph = await self.create_graph(parsed_papers=parsed_papers) - if self.debug: - logger.debug(f"Created {graph}") - selected_nodes = await select_nodes(graph, query=gap, llm=self.llm) - if self.debug: - logger.debug( - f"Selected nodes from {len(selected_nodes)} papers. Now generating answer", - ) - output, attributed_source_answers = await generate_final_answer( - graph, - query=gap, - all_selected_nodes=selected_nodes, - llm=self.llm, - ) - - graph = json_graph.node_link_data(graph, edges="edges") if self.config.output_graph else None - - return GapOutputSchema( - output=output, - attributed_source_answers=attributed_source_answers, - graph=graph, - ) - - async def _arun(self, params: GapInputSchema, **kwargs) -> GapOutputSchema: - return await self.get_response_async(params, **kwargs) diff --git a/akd/agents/gap_analysis/graph_utils.py b/akd/agents/gap_analysis/graph_utils.py deleted file mode 100644 index 1aae2744..00000000 --- a/akd/agents/gap_analysis/graph_utils.py +++ /dev/null @@ -1,446 +0,0 @@ -import ast -import asyncio -from typing import Dict, List, Tuple, Union - -import networkx as nx - -from .prompts import ( - GEN_ANSWER_PROMPT, - SECTION_CLASSIFIER_PROMPT, - SELECT_SUBSECTIONS_PROMPT, - SUMMARISE_ANSWER_PROMPT, - TRAVERSE_RELATIONS_PROMPT, -) -from .structures import PaperDataItem, ParsedPaper - - -async def add_paper_to_graph(graph: nx.Graph, paper: ParsedPaper, llm) -> nx.Graph: - """ - Adds a paper and its associated metadata to a NetworkX graph. - - This function creates a node for the paper and connects it with the following edges for now: - - Author nodes via `authored_by` edges - - Cited and referencing papers via `refers_to` and `cited_by` edges - - Section and subsection nodes if the paper is a `ParsedPaper`, using LLM classification for section types - - Args: - paper (ParsedPaper): The paper object to add to the graph. - G (networkx.Graph): The graph to update with paper data. - llm: The language model used for classifying section titles (only used for `ParsedPaper`). - - Returns: - networkx.Graph: The updated graph containing the new paper and its connections. - """ - paper_data = paper.model_dump(exclude_unset=False) - - paper_id = paper_data.get("paper_id") - if not paper_id: - raise ValueError("paper_id is required to add node to the graph.") - - if paper_id not in graph: - # Add the paper node with all fields - graph.add_node(paper_id, **paper_data, node_type="paper") - - # Add author nodes and edges - authors = paper_data.get("authors") or [] - for author in authors: - author_id = author.get("authorId") - author_name = author.get("name") - if author_id: - if author_id not in graph: - graph.add_node( - author_id, - name=author_name, - title=author_name, - node_type="author", - ) - graph.add_edge(paper_id, author_id, relationship="authored_by") - - # Handle citations and references - def add_related_papers(related_list, relationship, node_type): - if not related_list: - return - for related in related_list: - related_id = related.get("paperId") or related.get("paper_id") - related_title = related.get("title", "") - if related_id: - if related_id not in graph: - graph.add_node( - related_id, - paperTitle=related_title, - title=related_title, - node_type=node_type, - ) - graph.add_edge(paper_id, related_id, relationship=relationship) - - add_related_papers( - paper_data.get("references"), - relationship="refers_to", - node_type="reference", - ) - add_related_papers( - paper_data.get("citations"), - relationship="cited_by", - node_type="citation", - ) - - # Handle parsed paper sections (specific to ParsedPaper) - if isinstance(paper, ParsedPaper): - sections = paper_data.get("sections") or [] - section_to_key = await classify_section_titles(paper, llm) - for i, section in enumerate(sections, start=1): - section_title = section.get("title") - section_content = section.get("content") - section_subsections = section.get("subsections") - - if section_title not in section_to_key: - print( - f"Section {section_title} missing for {paper_data.get('title')}", - ) - continue - - section_node_type = section_to_key[section_title] - if section_node_type == "misc": - continue - - section_id = f"{paper_id}_{i}" - graph.add_node( - section_id, - section_title=section_title, - section_content=section_content, - node_type=section_node_type, - ) - graph.add_edge(paper_id, section_id, relationship="contains_section") - - if section_subsections: - for j, subsection in enumerate(section_subsections, start=1): - subsection_id = f"{paper_id}_{i}_{j}" - graph.add_node( - subsection_id, - subsection_title=subsection.get("title"), - subsection_content=subsection.get("content"), - node_type=section_node_type, - ) - graph.add_edge( - section_id, - subsection_id, - relationship="contains_subsection", - ) - - return graph - - -def get_nodes_by_type(graph: nx.Graph, node_type: str) -> List[str]: - """ - Retrieves all nodes from the graph that match a specific node type. - - Args: - G (networkx.Graph): The graph from which to retrieve nodes. - node_type (str): The type of node to filter by (e.g., 'paper', 'author', 'section'). - - Returns: - List[str]: A list of node identifiers matching the specified node type. - """ - nodes = [ - node - for node, attrs in graph.nodes(data=True) - if attrs.get("node_type") == node_type - ] - return nodes - - -def extract_direct_triples( - graph: nx.Graph, - start_node: str, -) -> List[Tuple[str, str, str]]: - """ - Extracts all direct relationship triples from the given start node in the graph. - - Args: - G (networkx.Graph): The graph containing nodes and edges. - start_node (str): The node ID from which to extract direct relationships. - - Returns: - List[Tuple[str, str, str]]: A deduplicated list of (node_type_1, relationship, node_type_2) triples. - """ - triples = [] - if start_node not in graph: - raise ValueError(f"Node {start_node} does not exist in the graph.") - for neighbor, data in graph[start_node].items(): - node_type_1 = graph.nodes[start_node].get("node_type", "unknown") - node_type_2 = graph.nodes[neighbor].get("node_type", "unknown") - relationship = data.get("relationship", "unknown") - triples.append((node_type_1, relationship, node_type_2)) - return list(set(triples)) - - -def get_connected_nodes( - graph: nx.Graph, - start_node: str, - target_node_type: str = None, - relation: str = None, - first_node: bool = False, -) -> Union[List, Dict, None]: - """ - Retrieves nodes directly connected to a given node, with optional filtering by node type and relationship. - - Args: - G (networkx.Graph): The graph containing nodes and edges. - start_node (str): The node ID from which to search for connected nodes. - target_node_type (Optional[str]): Filter to return only nodes of this type. If None, all types are considered. - relation (Optional[str]): Filter to return only edges with this relationship label. If None, all relationships are considered. - first_node (bool): If True, return only the first matching node's attributes (dict). If False, return all matches. - - Returns: - Union[List[Tuple[str, str, dict]], dict, None]: - - If first_node=False: A list of tuples (node_id, relationship, node_attributes) matching the criteria. - - If first_node=True: A single node's attributes (dict) if found, else None. - """ - if start_node not in graph: - raise ValueError(f"Node {start_node} does not exist in the graph.") - connected_nodes = [] - for neighbor, data in graph[start_node].items(): - neighbor_node_type = graph.nodes[neighbor].get("node_type", "unknown") - relationship = data.get("relationship", "unknown") - if (target_node_type is None or neighbor_node_type == target_node_type) and ( - relation is None or relation == relationship - ): - connected_nodes.append((neighbor, relationship, graph.nodes[neighbor])) - if first_node: - return graph.nodes[neighbor] - return connected_nodes if not first_node else None - - -async def select_subsections( - connected_subsection_nodes: List, - query: str, - llm, -) -> List[str]: - """ - Selects relevant subsection nodes from a list of connected nodes based on a user query using an LLM. - - Args: - connected_subsection_nodes (List[Tuple[str, str, dict]]): List of tuples representing subsection nodes - connected to a parent node. Each tuple contains (node_id, relation, node_data). - query (str): The user question or query. - llm: The language model instance used to evaluate relevance. - - Returns: - List[str]: A list of node IDs corresponding to the selected relevant subsections. - """ - title_to_node_id_map = {} - subsection_titles = [] - selected_subsection_nodes = [] - for node_id, relation, data in connected_subsection_nodes: - subsection_title = f"{data['node_type']} {relation} {data['subsection_title']}" - subsection_titles.append(subsection_title) - title_to_node_id_map[subsection_title] = node_id - if len(subsection_titles) != 0: - output = await llm.ainvoke( - SELECT_SUBSECTIONS_PROMPT.format_prompt( - query=query, - titles=subsection_titles, - ), - ) - selected_subsection_titles = ast.literal_eval(output.content) - selected_subsection_nodes = [ - title_to_node_id_map[subsection_title] - for subsection_title in selected_subsection_titles - ] - return selected_subsection_nodes - - -async def retrieve_relevant_sections( - graph: nx.Graph, - paper_node: str, - query: str, - llm, -) -> List[str]: - """ - Retrieves relevant section and subsection nodes from a graph based on a user query using an LLM. - - Args: - G (networkx.Graph): The graph containing paper and section nodes. - paper_node (str): The node ID of the paper in the graph. - query (str): The user query for retrieving relevant sections. - llm: The language model instance used for relevance evaluation. - - Returns: - List[str]: A list of node IDs representing relevant sections and subsections. - """ - relations = extract_direct_triples(graph, paper_node) - relation_traversal_output = await llm.ainvoke( - TRAVERSE_RELATIONS_PROMPT.invoke({"query": query, "relations": relations}), - ) - selected_sections = ast.literal_eval(relation_traversal_output.content) - connected_section_nodes = [] - for selected_section in selected_sections: - _, relation, target_node_type = selected_section - section_nodes = get_connected_nodes( - graph, - paper_node, - target_node_type, - relation, - ) - connected_section_nodes.extend(section_nodes) - - connected_subsection_nodes = [] - for section_node in connected_section_nodes: - start_node_id, relation, _ = section_node - subsection_nodes = get_connected_nodes( - graph, - start_node_id, - relation="contains_subsection", - ) - connected_subsection_nodes.extend(subsection_nodes) - - selected_subsection_nodes = await select_subsections( - connected_subsection_nodes, - query, - llm, - ) - - selected_nodes = [] - for conn_section_node in connected_section_nodes: - node_id, _, data = conn_section_node - if len(data["section_content"]) > 0: - selected_nodes.append(node_id) - selected_nodes.extend(selected_subsection_nodes) - return selected_nodes - - -async def classify_section_titles(paper: PaperDataItem, llm) -> Dict[str, str]: - """ - Classifies section titles of a paper into standardized categories using an LLM. - - Args: - paper (PaperDataItem): The paper object containing grouped section titles. - llm: Language model instance used for classification. - - Returns: - dict[str, str]: A mapping from each section title to its classified category label. - """ - section_titles = paper.section_titles - main_section_titles = [section_list[0] for section_list in section_titles] - classify_sections_chain = SECTION_CLASSIFIER_PROMPT | llm - out = await classify_sections_chain.ainvoke( - {"sections_to_group": main_section_titles}, - ) - grouped_section_titles = ast.literal_eval(out.content) - section_to_key = {} - for key, item in grouped_section_titles.items(): - for val in item: - section_to_key[val] = key - return section_to_key - - -async def select_nodes( - graph: nx.Graph, - query: str, - llm, -) -> List[List[str]]: - """ - Selects relevant section and subsection nodes across all paper nodes in the graph based on a query. - - Args: - G (networkx.Graph): The graph containing paper and related nodes. - query (str): The user query to find relevant sections. - llm: The language model instance used to evaluate relevance. - - Returns: - List[List[str]]: A list where each element is a list of node IDs relevant to a particular paper. - """ - paper_nodes = get_nodes_by_type(graph, "paper") - tasks = [ - retrieve_relevant_sections(graph, paper_node, query, llm) - for paper_node in paper_nodes - ] - all_selected_nodes = await asyncio.gather(*tasks) - return all_selected_nodes - - -def format_attributed_answers( - graph: nx.Graph, - attributed_answers: Dict[str, str], -) -> Dict[str, Dict[str, str]]: - """ - Formats attributed answers with metadata from the graph. - - Args: - G (networkx.Graph): The graph containing paper and section nodes. - attributed_answers (Dict[str, str]): A dictionary mapping source node IDs to generated content. - - Returns: - Dict[str, Dict[str, str]]: A dictionary where each key is a source ID and the value is a dictionary - containing the paper title, section title, URL, and the associated content. - """ - attributed_content = {} - for source_id, content in attributed_answers.items(): - section_title_id = source_id.strip().split(" ")[-1] - source_paper_id = section_title_id.split("_")[0] - title = graph.nodes[source_paper_id]["title"] - if section_title_id not in graph.nodes: - continue - section_data = graph.nodes[section_title_id] - if "section_title" in section_data.keys(): - section_title = section_data["section_title"] - else: - section_title = section_data["subsection_title"] - attributed_content[source_id] = { - "title": title, - "section_title": section_title, - "url": graph.nodes[source_paper_id]["url"], - "content": content, - } - return attributed_content - - -async def generate_final_answer( - graph: nx.Graph, - query: str, - all_selected_nodes: List[List[str]], - llm, -) -> str: - """ - Generates a consolidated answer for the query based on selected nodes in the graph. - - Args: - G (networkx.Graph): The graph containing paper and section nodes. - query (str): The user query for which the answer is generated. - all_selected_nodes (List[List[str]]): A list of lists of node IDs selected as relevant for the query. - llm: The language model instance used for generating and summarizing answers. - - Returns: - Response: The output of the summarization chain containing the final consolidated answer. - """ - local_answer_chain = GEN_ANSWER_PROMPT | llm - final_answer_chain = SUMMARISE_ANSWER_PROMPT | llm - all_answer_prompt_data = [] - for selected_nodes in all_selected_nodes: - answer_prompt_data = [] - for selected_node in selected_nodes: - node_data = graph.nodes[selected_node] - answer_prompt_data.append( - { - "query": query, - "section_title": node_data[list(node_data.keys())[0]], - "node_type": node_data["node_type"], - "section_content": node_data[list(node_data.keys())[1]], - }, - ) - all_answer_prompt_data.append(answer_prompt_data) - all_answers = await asyncio.gather( - *[ - local_answer_chain.abatch(answer_prompt_data) - for answer_prompt_data in all_answer_prompt_data - ], - ) - attributed_answers = {} - for answers, selected_nodes in zip(all_answers, all_selected_nodes): - for node_id, answer in zip(selected_nodes, answers): - attributed_answers.update({node_id: answer.content}) - output = await final_answer_chain.ainvoke( - input={"query": query, "attributed_answer_list": attributed_answers}, - ) - attributed_source_answers = format_attributed_answers(graph, attributed_answers) - return output.content, attributed_source_answers diff --git a/akd/agents/gap_analysis/parsing_utils.py b/akd/agents/gap_analysis/parsing_utils.py deleted file mode 100644 index 51879c09..00000000 --- a/akd/agents/gap_analysis/parsing_utils.py +++ /dev/null @@ -1,136 +0,0 @@ -import ast -from typing import Dict, List - -from bs4 import BeautifulSoup - -from .prompts import SECTION_GROUPER_PROMPT -from .structures import Section, SubSection - - -def parse_html(html_content: str) -> List[Dict[str, str]]: - """ - Parses the given HTML content to extract sections and their textual content. - - The function identifies all

headers as section titles. For each section, it collects the text - from sibling elements until the next

is encountered. Supported sibling elements include - paragraphs (

), lists (