diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 8f4dc66b86e..3b04ba06be6 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -61,6 +61,7 @@ /python/packages/claude/ @chetantoshniwal @eavanvalkenburg @giles17 @moonbox3 /python/packages/copilotstudio/ @chetantoshniwal @eavanvalkenburg @giles17 @moonbox3 /python/packages/core/ @chetantoshniwal @eavanvalkenburg @moonbox3 @TaoChenOSU @jpalvarezl @giles17 +/python/packages/core/agent_framework/_vectors.py @chetantoshniwal @westey-m @eavanvalkenburg @giles17 @moonbox3 @TaoChenOSU @jpalvarezl @peibekwe @baywet @rogerbarreto @SergeyMenshykh /python/packages/core/agent_framework/_workflows/ @chetantoshniwal @eavanvalkenburg @moonbox3 @TaoChenOSU @jpalvarezl /python/packages/core/agent_framework/_harness/ @chetantoshniwal @westey-m @eavanvalkenburg @moonbox3 /python/packages/declarative/ @chetantoshniwal @eavanvalkenburg @moonbox3 @peibekwe @baywet diff --git a/docs/features/vector-stores-and-embeddings/README.md b/docs/features/vector-stores-and-embeddings/README.md index 9f820ad7c79..15e058fd58a 100644 --- a/docs/features/vector-stores-and-embeddings/README.md +++ b/docs/features/vector-stores-and-embeddings/README.md @@ -10,7 +10,7 @@ This feature ports the vector store abstractions, embedding generator abstractio | Vector store collections | CRUD operations on vector store collections (upsert, get, delete) | | Vector search | Unified search interface with `search_type` parameter (`"vector"`, `"keyword_hybrid"`) | | Data model decorator | `@vectorstoremodel` decorator for defining vector store data models (supports Pydantic, dataclasses, plain classes, dicts) | -| Agent tools | `create_search_tool`, `create_upsert_tool`, `create_get_tool`, `create_delete_tool` for agent-usable vector store operations | +| Agent tools | `create_vector_search_tool`, `create_upsert_tool`, `create_get_tool`, `create_delete_tool` for agent-usable vector store operations | | In-memory store | Zero-dependency vector store for testing and development | | 13+ connectors | Azure AI Search, Qdrant, Redis, PostgreSQL, MongoDB, Cosmos DB, Pinecone, Chroma, Weaviate, Oracle, SQL Server, FAISS | @@ -47,12 +47,12 @@ This feature ports the vector store abstractions, embedding generator abstractio - **Embedding types** (`Embedding`, `GeneratedEmbeddings`, `EmbeddingGenerationOptions`) in `agent_framework/_types.py` - **Embedding protocol + base class** (`SupportsGetEmbeddings`, `BaseEmbeddingClient`) in `agent_framework/_clients.py` - **All vector store specific code** in a new `agent_framework/_vectors.py` module — this includes: - - Enums: `FieldTypes`, `IndexKind`, `DistanceFunction` + - String literal aliases: `FieldTypes`, `IndexKind`, `DistanceFunction` - `VectorStoreField`, `VectorStoreCollectionDefinition` - - `SearchOptions`, `SearchResponse`, `RecordFilterOptions` + - `SearchResponse`, `SearchResults`, and explicit CRUD/search keyword arguments - `@vectorstoremodel` decorator - - Serialization/deserialization protocols - - `VectorStoreRecordHandler`, `BaseVectorCollection`, `BaseVectorStore`, `BaseVectorSearch` + - `register_vectorstoremodel` with msgspec-backed default codecs and optional custom codecs + - Internal record conversion shared by `BaseVectorCollection` and `BaseVectorSearch` - `SupportsVectorUpsert`, `SupportsVectorSearch` protocols - **OpenAI embeddings** in `agent_framework/openai/` (built into core, like OpenAI chat) - **Azure OpenAI embeddings** in `agent_framework/azure/` (built into core, follows `AzureOpenAIChatClient` pattern) @@ -69,9 +69,9 @@ This feature ports the vector store abstractions, embedding generator abstractio | `VectorStoreCollection` | `BaseVectorCollection` | Drop redundant `Store`, add `Base` prefix per AF pattern | | `VectorStore` | `BaseVectorStore` | Add `Base` prefix per AF pattern | | `VectorSearch` | `BaseVectorSearch` | Add `Base` prefix per AF pattern | -| `VectorSearchOptions` | `SearchOptions` | Shorter — context is already vector search | +| `VectorSearchOptions` | Explicit `search()` keyword arguments | Avoid an options object that only forwards values | | `VectorSearchResult` | `SearchResponse` | Align with `ChatResponse`/`AgentResponse` | -| `GetFilteredRecordOptions` | `RecordFilterOptions` | Shorter, more natural | +| `GetFilteredRecordOptions` | Explicit `get()` keyword arguments | Avoid an options object that only forwards values | | `EmbeddingGeneratorBase` | `BaseEmbeddingClient` | Matches AF `BaseChatClient` pattern | | `VectorStoreCollectionProtocol` | `SupportsVectorUpsert` | AF `Supports*` naming convention | | `VectorSearchProtocol` | `SupportsVectorSearch` | AF `Supports*` naming convention | @@ -88,7 +88,6 @@ This feature ports the vector store abstractions, embedding generator abstractio | `@vectorstoremodel` | `_vectors.py` | | `VectorStoreField` | `_vectors.py` | | `VectorStoreCollectionDefinition` | `_vectors.py` | -| `VectorStoreRecordHandler` | `_vectors.py` | | `FieldTypes` | `_vectors.py` | | `IndexKind` | `_vectors.py` | | `DistanceFunction` | `_vectors.py` | @@ -107,7 +106,7 @@ This feature ports the vector store abstractions, embedding generator abstractio | `EmbeddingTelemetryLayer` | `observability.py` | MRO-based OTel tracing for embeddings | | `SupportsVectorUpsert` | `_vectors.py` | Protocol for collection CRUD | | `SupportsVectorSearch` | `_vectors.py` | Protocol for vector search | -| `create_search_tool` | `_vectors.py` | Creates AF `FunctionTool` from vector search | +| `create_vector_search_tool` | `_vectors.py` | Creates AF `FunctionTool` from vector search | ## Source Files Reference (SK → AF mapping) @@ -187,41 +186,52 @@ This feature ports the vector store abstractions, embedding generator abstractio ### Phase 3: Core Vector Store Abstractions **Goal:** Establish all vector store types, enums, the decorator, collection definition, and base classes. **Mergeable:** Yes — adds new abstractions, no breaking changes. +**Feature stage:** Experimental (`VECTOR_STORES`). -#### 3.1 — Vector store enums and field types in `_vectors.py` -- `FieldTypes` enum: `KEY`, `VECTOR`, `DATA` -- `IndexKind` enum: `HNSW`, `FLAT`, `IVF_FLAT`, `DISK_ANN`, `QUANTIZED_FLAT`, `DYNAMIC`, `DEFAULT` -- `DistanceFunction` enum: `COSINE_SIMILARITY`, `COSINE_DISTANCE`, `DOT_PROD`, `EUCLIDEAN_DISTANCE`, `EUCLIDEAN_SQUARED_DISTANCE`, `MANHATTAN`, `HAMMING`, `DEFAULT` -- No `SearchType` enum — use `Literal["vector", "keyword_hybrid"]` instead, per AF convention of avoiding unnecessary imports +#### 3.1 — Vector store literal aliases and field types in `_vectors.py` +- `FieldTypes`: `Literal["key", "vector", "data"]` +- `IndexKind`: literal alias covering `hnsw`, `flat`, `ivf_flat`, `disk_ann`, `quantized_flat`, `dynamic`, and `default` +- `DistanceFunction`: literal alias covering the supported similarity and distance functions +- `SearchType`: `Literal["vector", "keyword_hybrid"]` - `VectorStoreField` plain class (not Pydantic) - `VectorStoreCollectionDefinition` class (not Pydantic internally, but supports Pydantic models as input) -- `SearchOptions` plain class — includes `score_threshold: float | None` for filtering results by score (see note below) - `SearchResponse` generic class -- `RecordFilterOptions` plain class +- `SearchResults` generic result container +- Explicit keyword arguments on `get()` and `search()` instead of options classes - `DISTANCE_FUNCTION_DIRECTION_HELPER` dict #### 3.2 — `@vectorstoremodel` decorator - Port from SK, works with dataclasses, Pydantic models, plain classes, and dicts +- Plain classes can declare `VectorStoreField` metadata on annotated `__init__` parameters, matching `@tool` - Sets `__vectorstoremodel__` and `__vectorstoremodel_definition__` on the class - Remove SK-specific `kernel` prefix (`__kernel_vectorstoremodel__` → `__vectorstoremodel__`) -#### 3.3 — Serialization/deserialization protocols -- `SerializeMethodProtocol`, `ToDictFunctionProtocol`, `FromDictFunctionProtocol`, etc. -- Port the record handler logic but without Pydantic base class — use plain class or ABC +#### 3.3 — Registered model codecs +- `register_vectorstoremodel` registers one collection definition and encoder/decoder pair per model type +- `@vectorstoremodel` creates the definition and registers msgspec-backed default codecs +- Dictionary records provide their collection definition directly +- DataFrames and other row containers convert to sequences of row mappings before using the batch API +- Custom encoder and decoder callbacks can be overridden independently +- Array-like values such as NumPy arrays serialize through their `tolist()` method without a NumPy dependency; + supply a custom decoder that calls `numpy.array` or `numpy.asarray` when the model should restore an array #### 3.4 — Vector store base classes in `_vectors.py` -- `VectorStoreRecordHandler` — internal base class that handles serialization/deserialization between user data models and store-specific formats, plus embedding generation for vector fields. Both `BaseVectorCollection` and `BaseVectorSearch` extend this. -- `BaseVectorCollection(VectorStoreRecordHandler)` — base for collections +- `_VectorStoreRecordHandler` — private base class that handles record conversion and embedding generation +- `BaseVectorCollection` — base for collections - Uses `SupportsGetEmbeddings` instead of `EmbeddingGeneratorBase` - Not a Pydantic model — use `__init__` with explicit params - - `upsert`, `get`, `delete`, `ensure_collection_exists`, `collection_exists`, `ensure_collection_deleted` + - Batch-oriented `upsert`, `get`, and `delete` + - `upsert()` generates vector values by default and requires an embedding generator for every vector field; + pass `generate_vectors=False` to preserve supplied vector values + - CRUD `get()` excludes vectors by default; pass `include_vectors=True` when stored embeddings are needed + - `ensure_collection_exists`, `collection_exists`, `ensure_collection_deleted` - Async context manager support - `BaseVectorStore` — base for stores - `get_collection`, `list_collection_names`, `collection_exists`, `ensure_collection_deleted` - Async context manager support #### 3.5 — Vector search base class -- `BaseVectorSearch(VectorStoreRecordHandler)` — base for vector search +- `BaseVectorSearch` — base for vector search - Single `search(search_type=...)` method with `search_type: Literal["vector", "keyword_hybrid"]` parameter — no enum, just a literal - `_inner_search` abstract method for implementations - Filter building with lambda parser (AST-based) @@ -234,14 +244,17 @@ This feature ports the vector store abstractions, embedding generator abstractio - No protocol for `VectorStore` — it's a factory for collections, not a capability to duck-type against #### 3.7 — Exception types -- Add vector store exceptions under `IntegrationException` or create new branch -- `VectorStoreException`, `VectorStoreOperationException`, `VectorSearchException`, `VectorStoreModelException`, etc. +- Use `ValueError` and `TypeError` for invalid arguments, model definitions, and record conversion +- Use `NotImplementedError` for connector capabilities that are not supported +- Use the existing `IntegrationException` and `IntegrationInvalidResponseException` at connector boundaries -#### 3.8 — `create_search_tool` on `BaseVectorSearch` -- Method on `BaseVectorSearch` that creates an AF `FunctionTool` from the vector search +#### 3.8 — `create_vector_search_tool` +- Standalone factory that creates an AF `FunctionTool` from any `SupportsVectorSearch` implementation - Wraps the single `search()` method, passing `search_type` parameter -- Accepts: `name`, `description`, `search_type`, `top`, `skip`, `filter`, `string_mapper` -- The tool takes a query string, vectorizes it, searches, and returns results as strings +- Accepts: `name`, `description`, `approval_mode`, `search_type`, `parameters`, `top`, `skip`, `filter`, `filter_mapper`, `result_mapper` +- Defaults to `query`; a custom Pydantic model or JSON schema can expose `top`, `skip`, and additional filter fields +- Custom schemas must require a string `query`; exposed `top` and `skip` fields must declare finite maximum values +- The tool vectorizes the query, searches, and maps results to text or multimodal `Content` - Can also be a standalone factory function in `_vectors.py` #### 3.9 — Tests for all vector store abstractions @@ -326,7 +339,7 @@ Each connector follows the AF package structure: #### 8.1 — `create_upsert_tool` — tool for upserting records into a collection #### 8.2 — `create_get_tool` — tool for retrieving records by key - Key-based lookup only (by primary key), not a search tool -- Documentation must clearly distinguish this from `create_search_tool`: get_tool retrieves specific records by their known key, while search_tool performs similarity/filtered search across the collection +- Documentation must clearly distinguish this from `create_vector_search_tool`: get_tool retrieves specific records by their known key, while the search tool performs similarity/filtered search across the collection - Consider if this overlaps with filtered search and document when to use which #### 8.3 — `create_delete_tool` — tool for deleting records by key #### 8.4 — Tests and samples for CRUD tools @@ -349,7 +362,7 @@ Each connector follows the AF package structure: **Mergeable:** Yes — independent of vector stores. #### 10.1 — TextSearch base class and types -- `SearchOptions`, `SearchResponse`, `TextSearchResult` +- `SearchResponse`, `TextSearchResult`, and explicit search keyword arguments - `TextSearch` base class with `search()` method - `create_search_function()` for kernel integration (may need AF equivalent) @@ -361,7 +374,7 @@ Each connector follows the AF package structure: ## Key Considerations -1. **No Pydantic for internal classes**: All AF internal classes should use plain classes. Pydantic is only used for user-facing input validation (e.g., vector store data models). +1. **msgspec-backed conversion**: Use msgspec as the default serialization/deserialization path. Pydantic and plain classes remain supported user-model adapters. 2. **Protocol + Base class**: Follow AF's pattern of both a `Protocol` for duck-typing and a `Base` ABC for implementation, matching how `SupportsChatGetResponse` + `BaseChatClient` works. @@ -375,11 +388,11 @@ Each connector follows the AF package structure: 7. **Reusable data models**: The `@vectorstoremodel` decorator and `VectorStoreCollectionDefinition` should be agnostic enough to work with both SK and AF. The core types (`FieldTypes`, `IndexKind`, `DistanceFunction`, `VectorStoreField`) should be identical or easily mapped. -8. **`create_search_tool`**: The AF-native equivalent of SK's `create_search_function`. Instead of creating a `KernelFunction`, this creates an AF `FunctionTool` (via the `@tool` decorator pattern) from a vector search. This allows agents to use vector search as a tool during conversations. Design: - - `create_search_tool(name, description, search_type, ...)` → returns a `FunctionTool` that wraps `VectorSearch.search(search_type=...)` - - The tool accepts a query string, performs embedding + vector search, and returns results as strings - - Supports configurable string mappers, filter functions, top/skip defaults - - Lives in `_vectors.py` as a method on `BaseVectorSearch` and/or as a standalone factory function +8. **`create_vector_search_tool`**: The AF-native equivalent of SK's `create_search_function`. Instead of creating a `KernelFunction`, this creates an AF `FunctionTool` from any `SupportsVectorSearch` implementation. This allows agents to use vector search as a tool during conversations. Design: + - `create_vector_search_tool(search, name, description, search_type, ...)` returns a `FunctionTool` + - The tool accepts declared parameters, performs embedding + vector search, and returns text or multimodal content + - Defaults to `query`; custom parameters can expose `top`, `skip`, and additional fields for the filter mapper + - Lives in `_vectors.py` without expanding the structural search protocol 9. **CRUD tools**: A full set of create/read/update/delete tools for vector store collections, allowing agents to manage data in vector stores. Design: - `create_upsert_tool(...)` → tool for upserting records @@ -387,4 +400,4 @@ Each connector follows the AF package structure: - `create_delete_tool(...)` → tool for deleting records - These are separate from search and are placed in a later phase -10. **Score threshold filtering**: `SearchOptions` includes `score_threshold: float | None` to filter search results by relevance score (ref: [SK .NET PR #13501](https://github.com/microsoft/semantic-kernel/pull/13501)). The semantics depend on the distance function: for similarity functions (cosine similarity, dot product), results *below* the threshold are filtered out; for distance functions (cosine distance, euclidean), results *above* the threshold are filtered out. Use `DISTANCE_FUNCTION_DIRECTION_HELPER` to determine direction. Connectors should implement this natively where the database supports it, falling back to client-side post-filtering otherwise. +10. **Score threshold filtering**: `search(score_threshold=...)` filters results by relevance score (ref: [SK .NET PR #13501](https://github.com/microsoft/semantic-kernel/pull/13501)). The semantics depend on the distance function: for similarity functions (cosine similarity, dot product), results *below* the threshold are filtered out; for distance functions (cosine distance, euclidean), results *above* the threshold are filtered out. Use `DISTANCE_FUNCTION_DIRECTION_HELPER` to determine direction. Connectors should implement this natively where the database supports it, falling back to client-side post-filtering otherwise. diff --git a/docs/specs/feature-usage-bit-registry.md b/docs/specs/feature-usage-bit-registry.md index 6b89e4988fb..1e955610acf 100644 --- a/docs/specs/feature-usage-bit-registry.md +++ b/docs/specs/feature-usage-bit-registry.md @@ -137,7 +137,8 @@ only to approved first-party endpoints. | 16 | `core.mcp_skills_source` | MCP-backed skills | `agent_framework.MCPSkillsSource` | | 17 | `core.session_store` | Agent session store | `agent_framework.SessionStore` / `FileSessionStore` | | 18 | `core.agent_hooks` | Agent Hooks middleware | `agent_framework.create_agent_hooks_middleware` | -| 19–31 | _reserved_ | core growth | — | +| 19 | `core.vector_stores` | Vector store abstractions | `BaseVectorCollection` / `BaseVectorSearch` operations | +| 20–31 | _reserved_ | core growth | — | | 32 | `orchestration.sequential` | Sequential orchestration | `agent_framework_orchestrations.SequentialBuilder` | | 33 | `orchestration.concurrent` | Concurrent orchestration | `agent_framework_orchestrations.ConcurrentBuilder` | | 34 | `orchestration.group_chat` | Group-chat orchestration | `agent_framework_orchestrations.GroupChatBuilder` | diff --git a/python/packages/core/AGENTS.md b/python/packages/core/AGENTS.md index 288e07a2a6e..7182615705d 100644 --- a/python/packages/core/AGENTS.md +++ b/python/packages/core/AGENTS.md @@ -13,6 +13,7 @@ agent_framework/ ├── _clients.py # Chat client base classes and protocols ├── _types.py # Core types (Message, ChatResponse, Content, etc.) ├── _tools.py # Tool definitions and function invocation +├── _vectors.py # Vector store models, CRUD/search abstractions, and protocols ├── _middleware.py # Middleware system for request/response interception ├── _sessions.py # AgentSession and context provider abstractions ├── _skills.py # Agent Skills system (models, executors, provider) @@ -64,6 +65,19 @@ agent_framework/ - **`@tool`** decorator - Converts functions to tools - **`use_function_invocation()`** - Decorator to add automatic function calling to chat clients +### Vector stores (`_vectors.py`) + +The vector store API is experimental under the shared `VECTOR_STORES` feature ID. + +- **`@vectorstoremodel`** - Declares key, data, and vector fields on dataclasses, Pydantic models, and plain classes +- **`register_vectorstoremodel`** - Registers one definition and msgspec-backed codec pair per model type +- **`BaseVectorCollection`** - Base class for collection lifecycle and msgspec-backed record CRUD operations; + upserts generate embeddings by default and retrieval excludes vectors by default +- **`BaseVectorStore`** - Base class for stores that create collection clients +- **`BaseVectorSearch`** - Base class for vector and keyword-hybrid search +- **`create_vector_search_tool`** - Creates an agent tool from any `SupportsVectorSearch` implementation +- **`SupportsVectorUpsert`** / **`SupportsVectorSearch`** - Structural protocols for vector store capabilities + ### Middleware (`_middleware.py`) - **`AgentMiddleware`** - Intercepts agent `run()` calls diff --git a/python/packages/core/agent_framework/__init__.py b/python/packages/core/agent_framework/__init__.py index 77c56fa7559..241bcda3aae 100644 --- a/python/packages/core/agent_framework/__init__.py +++ b/python/packages/core/agent_framework/__init__.py @@ -284,6 +284,25 @@ "validate_tool_mode", "validate_tools", ), + "._vectors": ( + "DISTANCE_FUNCTION_DIRECTION_HELPER", + "BaseVectorCollection", + "BaseVectorSearch", + "BaseVectorStore", + "DistanceFunction", + "FieldTypes", + "IndexKind", + "SearchResponse", + "SearchResults", + "SearchType", + "SupportsVectorSearch", + "SupportsVectorUpsert", + "VectorStoreCollectionDefinition", + "VectorStoreField", + "create_vector_search_tool", + "register_vectorstoremodel", + "vectorstoremodel", + ), "._workflows._agent": ("WorkflowAgent",), "._workflows._agent_executor": ("AgentExecutor", "AgentExecutorRequest", "AgentExecutorResponse"), "._workflows._agent_utils": ("resolve_agent_id",), @@ -371,6 +390,7 @@ "DEFAULT_MODE_SOURCE_ID", "DEFAULT_TODO_SOURCE_ID", "DEFAULT_TOOL_APPROVAL_SOURCE_ID", + "DISTANCE_FUNCTION_DIRECTION_HELPER", "EXCLUDED_KEY", "EXCLUDE_REASON_KEY", "GROUP_ANNOTATION_KEY", @@ -412,6 +432,9 @@ "BaseAgent", "BaseChatClient", "BaseEmbeddingClient", + "BaseVectorCollection", + "BaseVectorSearch", + "BaseVectorStore", "CachingSkillsSource", "Case", "CharacterEstimatorTokenizer", @@ -438,6 +461,7 @@ "DeduplicatingSkillsSource", "Default", "DelegatingSkillsSource", + "DistanceFunction", "Edge", "EdgeCondition", "EdgeDuplicationError", @@ -456,6 +480,7 @@ "ExperimentalFeature", "FanInEdgeGroup", "FanOutEdgeGroup", + "FieldTypes", "FileAccessProvider", "FileCheckpointStorage", "FileHistoryProvider", @@ -490,6 +515,7 @@ "InMemoryHistoryProvider", "InMemorySkillsSource", "InProcRunnerContext", + "IndexKind", "InlineSkill", "InlineSkillResource", "InlineSkillScript", @@ -527,6 +553,9 @@ "Runner", "RunnerContext", "SamplingApprovalCallback", + "SearchResponse", + "SearchResults", + "SearchType", "SecretString", "SelectiveToolCallCompactionStrategy", "ServiceSessionId", @@ -555,6 +584,8 @@ "SupportsImageGenerationTool", "SupportsMCPTool", "SupportsShellTool", + "SupportsVectorSearch", + "SupportsVectorUpsert", "SupportsWebSearchTool", "SwitchCaseEdgeGroup", "SwitchCaseEdgeGroupCase", @@ -580,6 +611,8 @@ "UsageDetails", "UserInputRequiredException", "ValidationTypeEnum", + "VectorStoreCollectionDefinition", + "VectorStoreField", "Workflow", "WorkflowAgent", "WorkflowBuilder", @@ -613,6 +646,7 @@ "create_always_approve_tool_with_arguments_response", "create_edge_runner", "create_harness_agent", + "create_vector_search_tool", "detect_media_type_from_base64", "enqueue_messages", "evaluate_agent", @@ -636,6 +670,7 @@ "prepend_instructions_to_messages", "register_checkpoint_type", "register_state_type", + "register_vectorstoremodel", "resolve_agent_id", "response_handler", "set_agent_mode", @@ -650,6 +685,7 @@ "validate_tool_mode", "validate_tools", "validate_workflow_graph", + "vectorstoremodel", "workflow", ] diff --git a/python/packages/core/agent_framework/__init__.pyi b/python/packages/core/agent_framework/__init__.pyi index fa8f6a75ae6..5816c90600f 100644 --- a/python/packages/core/agent_framework/__init__.pyi +++ b/python/packages/core/agent_framework/__init__.pyi @@ -250,6 +250,25 @@ from ._types import ( validate_tool_mode, validate_tools, ) +from ._vectors import ( + DISTANCE_FUNCTION_DIRECTION_HELPER, + BaseVectorCollection, + BaseVectorSearch, + BaseVectorStore, + DistanceFunction, + FieldTypes, + IndexKind, + SearchResponse, + SearchResults, + SearchType, + SupportsVectorSearch, + SupportsVectorUpsert, + VectorStoreCollectionDefinition, + VectorStoreField, + create_vector_search_tool, + register_vectorstoremodel, + vectorstoremodel, +) from ._workflows._agent import WorkflowAgent from ._workflows._agent_executor import AgentExecutor, AgentExecutorRequest, AgentExecutorResponse from ._workflows._agent_utils import resolve_agent_id @@ -335,6 +354,7 @@ __all__ = [ "DEFAULT_MODE_SOURCE_ID", "DEFAULT_TODO_SOURCE_ID", "DEFAULT_TOOL_APPROVAL_SOURCE_ID", + "DISTANCE_FUNCTION_DIRECTION_HELPER", "EXCLUDED_KEY", "EXCLUDE_REASON_KEY", "GROUP_ANNOTATION_KEY", @@ -376,6 +396,9 @@ __all__ = [ "BaseAgent", "BaseChatClient", "BaseEmbeddingClient", + "BaseVectorCollection", + "BaseVectorSearch", + "BaseVectorStore", "CachingSkillsSource", "Case", "CharacterEstimatorTokenizer", @@ -402,6 +425,7 @@ __all__ = [ "DeduplicatingSkillsSource", "Default", "DelegatingSkillsSource", + "DistanceFunction", "Edge", "EdgeCondition", "EdgeDuplicationError", @@ -420,6 +444,7 @@ __all__ = [ "ExperimentalFeature", "FanInEdgeGroup", "FanOutEdgeGroup", + "FieldTypes", "FileAccessProvider", "FileCheckpointStorage", "FileHistoryProvider", @@ -454,6 +479,7 @@ __all__ = [ "InMemoryHistoryProvider", "InMemorySkillsSource", "InProcRunnerContext", + "IndexKind", "InlineSkill", "InlineSkillResource", "InlineSkillScript", @@ -491,6 +517,9 @@ __all__ = [ "Runner", "RunnerContext", "SamplingApprovalCallback", + "SearchResponse", + "SearchResults", + "SearchType", "SecretString", "SelectiveToolCallCompactionStrategy", "ServiceSessionId", @@ -519,6 +548,8 @@ __all__ = [ "SupportsImageGenerationTool", "SupportsMCPTool", "SupportsShellTool", + "SupportsVectorSearch", + "SupportsVectorUpsert", "SupportsWebSearchTool", "SwitchCaseEdgeGroup", "SwitchCaseEdgeGroupCase", @@ -544,6 +575,8 @@ __all__ = [ "UsageDetails", "UserInputRequiredException", "ValidationTypeEnum", + "VectorStoreCollectionDefinition", + "VectorStoreField", "Workflow", "WorkflowAgent", "WorkflowBuilder", @@ -577,6 +610,7 @@ __all__ = [ "create_always_approve_tool_with_arguments_response", "create_edge_runner", "create_harness_agent", + "create_vector_search_tool", "detect_media_type_from_base64", "enqueue_messages", "evaluate_agent", @@ -600,6 +634,7 @@ __all__ = [ "prepend_instructions_to_messages", "register_checkpoint_type", "register_state_type", + "register_vectorstoremodel", "resolve_agent_id", "response_handler", "set_agent_mode", @@ -614,5 +649,6 @@ __all__ = [ "validate_tool_mode", "validate_tools", "validate_workflow_graph", + "vectorstoremodel", "workflow", ] diff --git a/python/packages/core/agent_framework/_feature_stage.py b/python/packages/core/agent_framework/_feature_stage.py index 53893e397f3..9f694f6423a 100644 --- a/python/packages/core/agent_framework/_feature_stage.py +++ b/python/packages/core/agent_framework/_feature_stage.py @@ -64,6 +64,7 @@ class ExperimentalFeature(str, Enum): PROGRESSIVE_TOOLS = "PROGRESSIVE_TOOLS" SESSION_STORE = "SESSION_STORE" TO_PROMPT_AGENT = "TO_PROMPT_AGENT" + VECTOR_STORES = "VECTOR_STORES" class ReleaseCandidateFeature(str, Enum): diff --git a/python/packages/core/agent_framework/_telemetry.py b/python/packages/core/agent_framework/_telemetry.py index 955e276fe2d..d20f31b86cf 100644 --- a/python/packages/core/agent_framework/_telemetry.py +++ b/python/packages/core/agent_framework/_telemetry.py @@ -55,6 +55,7 @@ class FeatureIndex(IntEnum): CORE_MCP_SKILLS_SOURCE = 16 CORE_SESSION_STORE = 17 CORE_AGENT_HOOKS = 18 + CORE_VECTOR_STORES = 19 # This environment variable is reserved by the Foundry hosting environment to diff --git a/python/packages/core/agent_framework/_vectors.py b/python/packages/core/agent_framework/_vectors.py new file mode 100644 index 00000000000..761cc3a5bb2 --- /dev/null +++ b/python/packages/core/agent_framework/_vectors.py @@ -0,0 +1,1923 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Core vector store abstractions.""" + +from __future__ import annotations + +import operator +from abc import ABC, abstractmethod +from ast import AST, Lambda, NodeVisitor, expr, parse +from collections.abc import AsyncIterable, AsyncIterator, Callable, Mapping, Sequence +from dataclasses import dataclass, is_dataclass, replace +from inspect import Parameter, getsource, signature +from types import UnionType +from typing import ( + Annotated, + Any, + ClassVar, + Final, + Generic, + Literal, + Protocol, + TypeAlias, + TypeGuard, + Union, + cast, + get_args, + get_origin, + get_type_hints, + overload, + runtime_checkable, +) + +import msgspec +from pydantic import BaseModel +from typing_extensions import Self, TypedDict, TypeVar + +from ._clients import SupportsGetEmbeddings +from ._feature_stage import ExperimentalFeature, experimental +from ._telemetry import FeatureIndex, mark_feature_used +from ._tools import FunctionTool +from ._types import Content, EmbeddingGenerationOptions +from .exceptions import IntegrationException, IntegrationInvalidResponseException + +ModelT = TypeVar("ModelT", default=Any) +KeyT = TypeVar("KeyT", default=Any) +FilterT = TypeVar("FilterT") +ResultT = TypeVar("ResultT") +DecoratedModelT = TypeVar("DecoratedModelT") + +SearchType: TypeAlias = Literal["vector", "keyword_hybrid"] +FieldTypes: TypeAlias = Literal["key", "vector", "data"] +IndexKind: TypeAlias = Literal["hnsw", "flat", "ivf_flat", "disk_ann", "quantized_flat", "dynamic", "default"] +DistanceFunction: TypeAlias = Literal[ + "cosine_similarity", + "cosine_distance", + "dot_prod", + "euclidean_distance", + "euclidean_squared_distance", + "manhattan", + "hamming", + "DEFAULT", +] +Vector: TypeAlias = Sequence[float | int] +RecordFilter: TypeAlias = Callable[[Any], bool] | str +RecordFilters: TypeAlias = RecordFilter | Sequence[RecordFilter] +EmbeddingClient: TypeAlias = SupportsGetEmbeddings[Any, Any, Any] +VectorModelEncoder: TypeAlias = Callable[[Any], Mapping[str, Any]] +VectorModelDecoder: TypeAlias = Callable[[Mapping[str, Any]], Any] + +_DEFAULT_SEARCH_TOOL_NAME: Final[str] = "search" +_DEFAULT_SEARCH_TOOL_DESCRIPTION: Final[str] = ( + "Perform a vector search for data in a vector store using the provided query." +) +_INDEX_KINDS: Final[tuple[str, ...]] = ( + "hnsw", + "flat", + "ivf_flat", + "disk_ann", + "quantized_flat", + "dynamic", + "default", +) +_DISTANCE_FUNCTIONS: Final[tuple[str, ...]] = ( + "cosine_similarity", + "cosine_distance", + "dot_prod", + "euclidean_distance", + "euclidean_squared_distance", + "manhattan", + "hamming", + "DEFAULT", +) + + +DISTANCE_FUNCTION_DIRECTION_HELPER: Final[Mapping[DistanceFunction, Callable[[float | int, float | int], bool]]] = { + "cosine_similarity": operator.ge, + "cosine_distance": operator.le, + "dot_prod": operator.ge, + "euclidean_distance": operator.le, + "euclidean_squared_distance": operator.le, + "manhattan": operator.le, + "hamming": operator.le, +} + + +def _msgspec_enc_hook(value: Any) -> Any: + if isinstance(value, BaseModel): + return value.model_dump() + to_list = getattr(value, "tolist", None) + if callable(to_list): + return to_list() + if hasattr(value, "__dict__"): + return cast(dict[str, Any], vars(value)) + raise NotImplementedError(f"Objects of type {type(value).__name__!r} are not supported.") + + +def _normalize_vector(value: Any) -> Vector: + if isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)): + return cast(Vector, value) + to_list = getattr(value, "tolist", None) + if callable(to_list): + converted = to_list() + if isinstance(converted, Sequence) and not isinstance(converted, (str, bytes, bytearray)): + return cast(Vector, converted) + raise TypeError("The embedding client returned an unsupported vector type.") + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +@dataclass(frozen=True, slots=True, init=False) +class VectorStoreField: + """Describe one field in a vector store model.""" + + field_type: FieldTypes + name: str + type_: str | None + storage_name: str | None + is_indexed: bool | None + is_full_text_indexed: bool | None + dimensions: int | None + index_kind: IndexKind | None + distance_function: DistanceFunction | None + embedding_generator: EmbeddingClient | None + + @overload + def __init__( + self, + field_type: Literal["key"], + *, + name: str | None = None, + type_: str | None = None, + storage_name: str | None = None, + ) -> None: + """Initialize a key field. + + Args: + field_type: The key field type. + name: The model field name. The decorator supplies this when omitted. + type_: The scalar type name used by the backing store. + storage_name: The field name used by the backing store. + """ + ... + + @overload + def __init__( + self, + field_type: Literal["data"] = "data", + *, + name: str | None = None, + type_: str | None = None, + storage_name: str | None = None, + is_indexed: bool | None = None, + is_full_text_indexed: bool | None = None, + ) -> None: + """Initialize a data field with optional indexing. + + Args: + field_type: The data field type. + name: The model field name. The decorator supplies this when omitted. + type_: The scalar type name used by the backing store. + storage_name: The field name used by the backing store. + is_indexed: Whether the field should be indexed. + is_full_text_indexed: Whether the field should have a full-text index. + """ + ... + + @overload + def __init__( + self, + field_type: Literal["vector"], + *, + name: str | None = None, + type_: str | None = None, + storage_name: str | None = None, + dimensions: int, + index_kind: IndexKind | None = None, + distance_function: DistanceFunction | None = None, + embedding_generator: EmbeddingClient | None = None, + ) -> None: + """Initialize a vector field with required dimensions. + + Args: + field_type: The vector field type. + name: The model field name. The decorator supplies this when omitted. + type_: The vector element type name used by the backing store. + storage_name: The field name used by the backing store. + dimensions: The number of vector dimensions. + index_kind: The vector index kind. + distance_function: The vector distance function. + embedding_generator: An optional client used to generate this field's embeddings. + + Raises: + ValueError: If dimensions or vector options are invalid. + """ + ... + + def __init__( + self, + field_type: FieldTypes = "data", + *, + name: str | None = None, + type_: str | None = None, + storage_name: str | None = None, + is_indexed: bool | None = None, + is_full_text_indexed: bool | None = None, + dimensions: int | None = None, + index_kind: IndexKind | None = None, + distance_function: DistanceFunction | None = None, + embedding_generator: EmbeddingClient | None = None, + ) -> None: + """Initialize a vector store field. + + Args: + field_type: The field's role in the vector store model. + name: The model field name. The decorator supplies this when omitted. + type_: The scalar type name used by the backing store. + storage_name: The field name used by the backing store. + is_indexed: Whether a data field should be indexed. + is_full_text_indexed: Whether a data field should have a full-text index. + dimensions: The number of vector dimensions. Required for vector fields. + index_kind: The vector index kind. + distance_function: The vector distance function. + embedding_generator: An optional client used to generate this field's embeddings. + + Raises: + ValueError: If field options are invalid. + """ + if field_type not in ("key", "vector", "data"): + raise ValueError(f"Unknown vector store field type '{field_type}'.") + resolved_dimensions: int | None = None + resolved_index_kind: IndexKind | None = None + resolved_distance_function: DistanceFunction | None = None + resolved_embedding_generator: EmbeddingClient | None = None + if field_type == "vector": + if dimensions is None or dimensions <= 0: + raise ValueError("Vector fields must specify a positive number of dimensions.") + if index_kind is not None and index_kind not in _INDEX_KINDS: + raise ValueError(f"Unknown vector index kind '{index_kind}'.") + if distance_function is not None and distance_function not in _DISTANCE_FUNCTIONS: + raise ValueError(f"Unknown vector distance function '{distance_function}'.") + resolved_dimensions = dimensions + resolved_index_kind = index_kind or "default" + resolved_distance_function = distance_function or "DEFAULT" + resolved_embedding_generator = embedding_generator + elif any(value is not None for value in (dimensions, index_kind, distance_function, embedding_generator)): + raise ValueError("Vector-only options can only be set on vector fields.") + + object.__setattr__(self, "field_type", field_type) + object.__setattr__(self, "name", name or "") + object.__setattr__(self, "type_", type_) + object.__setattr__(self, "storage_name", storage_name) + object.__setattr__(self, "is_indexed", is_indexed) + object.__setattr__(self, "is_full_text_indexed", is_full_text_indexed) + object.__setattr__(self, "dimensions", resolved_dimensions) + object.__setattr__(self, "index_kind", resolved_index_kind) + object.__setattr__(self, "distance_function", resolved_distance_function) + object.__setattr__(self, "embedding_generator", resolved_embedding_generator) + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +@dataclass(frozen=True, slots=True, init=False) +class VectorStoreCollectionDefinition: + """Describe the records stored in a vector collection. + + Most users should not create this class directly. Applying + :func:`vectorstoremodel` to a typed model derives and registers its + collection definition automatically. + + Create a definition explicitly for schema-less records such as dictionaries, + or when adapting an externally owned model through + :func:`register_vectorstoremodel`. + """ + + fields: tuple[VectorStoreField, ...] + collection_name: str | None + key_name: str + + def __init__( + self, + fields: Sequence[VectorStoreField], + *, + collection_name: str | None = None, + ) -> None: + """Initialize a vector store collection definition. + + Args: + fields: The key, data, and vector fields in each record. + collection_name: The collection name associated with the model. + + Raises: + ValueError: If field names or key fields are invalid. + """ + object.__setattr__(self, "fields", tuple(fields)) + object.__setattr__(self, "collection_name", collection_name) + object.__setattr__(self, "key_name", self._validate()) + + def _validate(self) -> str: + if not self.fields: + raise ValueError("A vector store definition must contain at least one field.") + if any(not field.name for field in self.fields): + raise ValueError("Vector store field names must not be empty.") + + names = [field.name for field in self.fields] + if len(names) != len(set(names)): + raise ValueError("Vector store field names must be unique.") + storage_names = [field.storage_name or field.name for field in self.fields] + if len(storage_names) != len(set(storage_names)): + raise ValueError("Vector store field storage names must be unique.") + + key_fields = [field for field in self.fields if field.field_type == "key"] + if len(key_fields) != 1: + raise ValueError("A vector store definition must contain exactly one key field.") + return key_fields[0].name + + @property + def names(self) -> list[str]: + """Get the model field names.""" + return [field.name for field in self.fields] + + @property + def storage_names(self) -> list[str]: + """Get the backing store field names.""" + return [field.storage_name or field.name for field in self.fields] + + @property + def key_field(self) -> VectorStoreField: + """Get the key field.""" + return next(field for field in self.fields if field.field_type == "key") + + @property + def key_field_storage_name(self) -> str: + """Get the key field's backing store name.""" + return self.key_field.storage_name or self.key_field.name + + @property + def vector_fields(self) -> list[VectorStoreField]: + """Get the vector fields.""" + return [field for field in self.fields if field.field_type == "vector"] + + @property + def data_fields(self) -> list[VectorStoreField]: + """Get the data fields.""" + return [field for field in self.fields if field.field_type == "data"] + + @property + def vector_field_names(self) -> list[str]: + """Get the vector field names.""" + return [field.name for field in self.vector_fields] + + @property + def data_field_names(self) -> list[str]: + """Get the data field names.""" + return [field.name for field in self.data_fields] + + def try_get_vector_field(self, field_name: str | None = None) -> VectorStoreField | None: + """Get a vector field by model or storage name, defaulting to the first vector field.""" + if field_name is None: + return self.vector_fields[0] if self.vector_fields else None + return next( + (field for field in self.vector_fields if field.name == field_name or field.storage_name == field_name), + None, + ) + + def get_names(self, *, include_vector_fields: bool = True, include_key_field: bool = True) -> list[str]: + """Get selected model field names.""" + return [ + field.name + for field in self.fields + if field.field_type == "data" + or (field.field_type == "vector" and include_vector_fields) + or (field.field_type == "key" and include_key_field) + ] + + def get_storage_names(self, *, include_vector_fields: bool = True, include_key_field: bool = True) -> list[str]: + """Get selected backing store field names.""" + return [ + field.storage_name or field.name + for field in self.fields + if field.field_type == "data" + or (field.field_type == "vector" and include_vector_fields) + or (field.field_type == "key" and include_key_field) + ] + + +@dataclass(frozen=True, slots=True) +class _VectorModelRegistration: + record_type: type[Any] + definition: VectorStoreCollectionDefinition + encoder: VectorModelEncoder + decoder: VectorModelDecoder + + +_VECTOR_MODEL_REGISTRY: dict[type[Any], _VectorModelRegistration] = {} + + +def _default_vector_model_encoder(record_type: type[Any]) -> VectorModelEncoder: + def encode(value: Any) -> Mapping[str, Any]: + if not isinstance(value, record_type): + raise TypeError(f"Expected {record_type.__name__}, got {type(value).__name__}.") + converted = msgspec.to_builtins(value, str_keys=True, enc_hook=_msgspec_enc_hook) + if not isinstance(converted, Mapping): + raise TypeError(f"Vector model {record_type.__name__!r} must serialize to a mapping.") + return cast(Mapping[str, Any], converted) + + return encode + + +def _default_vector_model_decoder(record_type: type[Any]) -> VectorModelDecoder: + if issubclass(record_type, BaseModel): + + def decode_pydantic(value: Mapping[str, Any]) -> Any: + validation_value = { + field.validation_alias + if isinstance(field.validation_alias, str) + else field.alias + if isinstance(field.alias, str) + else name: value[name] + for name, field in record_type.model_fields.items() + if name in value + } + return record_type.model_validate(validation_value) + + return decode_pydantic + if is_dataclass(record_type) or issubclass(record_type, msgspec.Struct): + return lambda value: msgspec.convert(value, record_type) + return lambda value: record_type(**value) + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +def register_vectorstoremodel( + record_type: type[ModelT], + *, + definition: VectorStoreCollectionDefinition, + encoder: Callable[[ModelT], Mapping[str, Any]] | None = None, + decoder: Callable[[Mapping[str, Any]], ModelT] | None = None, +) -> None: + """Register one vector store definition and codec pair for a model type. + + Args: + record_type: The model type to register. + definition: The vector store collection definition for the model. + encoder: Optional callback that converts a model instance to a mapping. + decoder: Optional callback that reconstructs a model instance from a mapping. + This can restore array-like fields such as NumPy arrays without requiring + Agent Framework to depend on NumPy. + + Raises: + ValueError: If the model type is already registered differently. + """ + existing = _VECTOR_MODEL_REGISTRY.get(record_type) + if existing is not None: + if existing.definition is not definition: + raise ValueError(f"Vector model {record_type.__name__!r} is already registered with another definition.") + if encoder is not None and existing.encoder is not encoder: + raise ValueError(f"Vector model {record_type.__name__!r} is already registered with another encoder.") + if decoder is not None and existing.decoder is not decoder: + raise ValueError(f"Vector model {record_type.__name__!r} is already registered with another decoder.") + return + if decoder is None: + required_vector_fields = [ + field.name for field in definition.vector_fields if not _has_default(record_type, field.name) + ] + if required_vector_fields: + raise ValueError( + "Vector fields omitted by include_vectors=False must declare defaults when using the default decoder. " + f"Add defaults or supply a custom decoder for: {', '.join(required_vector_fields)}." + ) + resolved_encoder = ( + cast(VectorModelEncoder, encoder) if encoder is not None else _default_vector_model_encoder(record_type) + ) + resolved_decoder = ( + cast(VectorModelDecoder, decoder) if decoder is not None else _default_vector_model_decoder(record_type) + ) + registration = _VectorModelRegistration( + record_type=record_type, + definition=definition, + encoder=resolved_encoder, + decoder=resolved_decoder, + ) + _VECTOR_MODEL_REGISTRY[record_type] = registration + + +def _has_default(record_type: type[Any], field_name: str) -> bool: + if issubclass(record_type, BaseModel) and field_name in record_type.model_fields: + return not record_type.model_fields[field_name].is_required() + try: + parameter = signature(record_type).parameters.get(field_name) + except (TypeError, ValueError): + parameter = None + if parameter is not None: + return parameter.default is not Parameter.empty + return hasattr(record_type, field_name) + + +def _unwrap_annotation(annotation: Any) -> Any: + if get_origin(annotation) is Annotated: + return get_args(annotation)[0] + return annotation + + +def _without_none(annotation: Any) -> tuple[Any, ...]: + args = get_args(annotation) + if get_origin(annotation) in (UnionType, Union): + return tuple(arg for arg in args if arg is not type(None)) + return (annotation,) + + +def _infer_type_name(annotation: Any, *, vector: bool) -> str | None: + candidates = _without_none(_unwrap_annotation(annotation)) + if vector: + for candidate in candidates: + origin = get_origin(candidate) + args = get_args(candidate) + if origin is not None and args: + candidate = next((arg for arg in args if arg is not Ellipsis), candidate) + return getattr(candidate, "__name__", str(candidate)) + candidate = candidates[0] if candidates else annotation + origin = get_origin(candidate) + return getattr(origin or candidate, "__name__", None) + + +def _parse_model_definition( + record_type: type[Any], + *, + collection_name: str | None, +) -> VectorStoreCollectionDefinition: + try: + annotations = get_type_hints(record_type, include_extras=True) + except (NameError, TypeError) as exc: + raise ValueError(f"Unable to resolve annotations for {record_type.__name__}: {exc}") from exc + uses_init_annotations = not any( + any(isinstance(metadata, VectorStoreField) for metadata in get_args(annotation)[1:]) + for annotation in annotations.values() + if get_origin(annotation) is Annotated + ) + init_parameters: Mapping[str, Parameter] = {} + if uses_init_annotations: + try: + annotations = { + name: annotation + for name, annotation in get_type_hints(record_type.__init__, include_extras=True).items() + if name not in {"self", "return"} + } + init_parameters = signature(record_type.__init__).parameters + except (NameError, TypeError, ValueError) as exc: + raise ValueError(f"Unable to resolve constructor annotations for {record_type.__name__}: {exc}") from exc + if not annotations: + raise ValueError("A vector store model must declare at least one annotated field or constructor parameter.") + + fields: list[VectorStoreField] = [] + for name, annotation in annotations.items(): + metadata = get_args(annotation)[1:] if get_origin(annotation) is Annotated else () + field = next((item for item in metadata if isinstance(item, VectorStoreField)), None) + if field is None: + has_default = ( + init_parameters[name].default is not Parameter.empty + if uses_init_annotations + else _has_default(record_type, name) + ) + if not has_default: + raise ValueError(f"Field '{name}' must use VectorStoreField metadata or declare a default value.") + continue + + parsed_field = replace( + field, + name=name, + type_=field.type_ or _infer_type_name(annotation, vector=field.field_type == "vector"), + ) + fields.append(parsed_field) + return VectorStoreCollectionDefinition(fields, collection_name=collection_name) + + +class _VectorStoreModelDecorator(Protocol): + def __call__(self, record_type: type[DecoratedModelT]) -> type[DecoratedModelT]: + """Decorate a model while preserving its concrete type.""" + ... + + +@overload +def vectorstoremodel(cls: type[ModelT]) -> type[ModelT]: + """Decorate a vector store model without arguments. + + Args: + cls: The class to decorate. + + Returns: + The original class with vector store model metadata attached. + + Raises: + ValueError: If the model definition is invalid. + """ + ... + + +@overload +def vectorstoremodel( + cls: None = None, + *, + collection_name: str | None = None, + encoder: Callable[[Any], Mapping[str, Any]] | None = None, + decoder: Callable[[Mapping[str, Any]], Any] | None = None, +) -> _VectorStoreModelDecorator: + """Create a vector store model decorator with a collection name. + + Args: + cls: The empty decorator target used when calling the decorator with arguments. + collection_name: The collection name associated with the model. + encoder: Optional callback that converts a model instance to a mapping. + decoder: Optional callback that reconstructs a model instance from a mapping. + This can restore array-like fields such as NumPy arrays without requiring + Agent Framework to depend on NumPy. + + Returns: + A decorator that attaches vector store model metadata. + + Raises: + ValueError: When the returned decorator receives an invalid model definition. + """ + ... + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +def vectorstoremodel( + cls: type[Any] | None = None, + *, + collection_name: str | None = None, + encoder: Callable[[Any], Mapping[str, Any]] | None = None, + decoder: Callable[[Mapping[str, Any]], Any] | None = None, +) -> type[Any] | _VectorStoreModelDecorator: + """Mark a class as a vector store model. + + Class fields or constructor parameters use ``Annotated`` metadata to describe their + vector store role. Dataclasses, Pydantic models, and plain classes are supported. + Dictionaries use an explicit :class:`VectorStoreCollectionDefinition` instead. + + Args: + cls: The class to decorate. + collection_name: The collection name associated with the model. + encoder: Optional callback that converts a model instance to a mapping. + decoder: Optional callback that reconstructs a model instance from a mapping. + + Returns: + The original class with vector store model metadata attached. + + Raises: + ValueError: If the model definition is invalid. + """ + + def wrap(record_type: type[DecoratedModelT]) -> type[DecoratedModelT]: + definition = _parse_model_definition(record_type, collection_name=collection_name) + register_vectorstoremodel( + record_type, + definition=definition, + encoder=encoder, + decoder=decoder, + ) + decorated_type = cast(Any, record_type) + decorated_type.__vectorstoremodel__ = True + decorated_type.__vectorstoremodel_definition__ = definition + return record_type + + return wrap if cls is None else wrap(cls) + + +def _validate_paging(*, top: int, skip: int) -> None: + if not isinstance(top, int) or isinstance(top, bool): + raise TypeError("top must be an integer.") + if not isinstance(skip, int) or isinstance(skip, bool): + raise TypeError("skip must be an integer.") + if top <= 0: + raise ValueError("top must be greater than zero.") + if skip < 0: + raise ValueError("skip must not be negative.") + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class SearchResponse(TypedDict, Generic[ModelT]): + """One vector search result.""" + + record: ModelT + score: float | None + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class SearchResults(Generic[ResultT]): + """A lazily consumed set of vector search results. + + Connector-native counts may be placed in ``metadata`` together with enough + provider-specific context to explain their scope. + """ + + def __init__( + self, + results: AsyncIterable[ResultT] | Sequence[ResultT], + *, + metadata: Mapping[str, Any] | None = None, + ) -> None: + """Initialize search results.""" + self.results = _as_async_iterable(results) + self.metadata = metadata + + def __aiter__(self) -> AsyncIterator[ResultT]: + """Iterate over results regardless of whether their source was synchronous or asynchronous.""" + return self.results.__aiter__() + + +class _VectorStoreRecordHandler(Generic[KeyT, ModelT]): + """Serialize and deserialize application records for a vector store.""" + + supported_key_types: ClassVar[set[str] | None] = None + supported_vector_types: ClassVar[set[str] | None] = None + + def __init__( + self, + record_type: type[ModelT], + *, + definition: VectorStoreCollectionDefinition | None = None, + embedding_generator: EmbeddingClient | None = None, + ) -> None: + """Initialize a vector store record handler. + + Args: + record_type: The application record type. + definition: The collection definition. Decorated models supply this automatically. + embedding_generator: The default client used for local vector generation. + + Raises: + ValueError: If no model registration or explicit dictionary definition is available. + """ + registration = _VECTOR_MODEL_REGISTRY.get(record_type) + if record_type is dict: + if definition is None: + raise ValueError("Dictionary record types require an explicit VectorStoreCollectionDefinition.") + resolved_definition = definition + else: + if registration is None: + raise ValueError( + f"Record type {record_type.__name__!r} must be registered with " + "@vectorstoremodel or register_vectorstoremodel()." + ) + if definition is not None and definition is not registration.definition: + raise ValueError(f"Record type {record_type.__name__!r} is registered with another definition.") + resolved_definition = registration.definition + self.record_type = record_type + self.definition = resolved_definition + self._model_registration = registration + self.embedding_generator = embedding_generator + self._validate_data_model() + + def _validate_data_model(self) -> None: + key_type = self.definition.key_field.type_ + if self.supported_key_types and key_type and key_type not in self.supported_key_types: + raise ValueError(f"Key field type must be one of {self.supported_key_types}; got '{key_type}'.") + if not self.supported_vector_types: + return + for field in self.definition.vector_fields: + if field.type_ and field.type_ not in self.supported_vector_types: + raise ValueError( + f"Vector field '{field.name}' type must be one of {self.supported_vector_types}; " + f"got '{field.type_}'." + ) + + def _serialize_dicts_to_store_models( + self, + records: Sequence[dict[str, Any]], + *, + context: Mapping[str, Any] | None = None, + ) -> Sequence[Any]: + """Convert dictionaries to store-specific records.""" + return records + + def _deserialize_store_models_to_dicts( + self, + records: Sequence[Any], + *, + context: Mapping[str, Any] | None = None, + ) -> Sequence[dict[str, Any]]: + """Convert store-specific records to dictionaries.""" + dict_records: list[dict[str, Any]] = [] + for record in records: + if not isinstance(record, Mapping): + raise TypeError("Store records must be mappings unless the collection overrides deserialization.") + dict_records.append(dict(cast(Mapping[str, Any], record))) + return dict_records + + async def serialize( + self, + records: ModelT | Sequence[ModelT], + *, + generate_vectors: bool = True, + context: Mapping[str, Any] | None = None, + ) -> Any: + """Serialize one or more application records for the backing store. + + Args: + records: One application record or a sequence of records. + generate_vectors: Whether to generate vector values, overwriting any supplied values. When ``False``, + supplied values are preserved. + context: Connector-specific serialization context. + + Raises: + TypeError: If a record cannot be converted to a mapping. + ValueError: If required record data is missing, has an invalid shape, or a vector field has no generator. + IntegrationInvalidResponseException: If embedding generation returns an unexpected result count. + """ + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + is_batch = _is_non_string_sequence(records) + input_records = list(cast(Sequence[ModelT], records)) if is_batch else [cast(ModelT, records)] + dict_records = [self._serialize_record_to_dict(record) for record in input_records] + + if generate_vectors: + await self._add_vectors_to_records(dict_records) + store_models = list(self._serialize_dicts_to_store_models(dict_records, context=context)) + + if len(store_models) != len(dict_records): + raise IntegrationInvalidResponseException( + f"Expected {len(dict_records)} serialized records, but the connector returned {len(store_models)}." + ) + if is_batch: + return store_models + if len(store_models) != 1: + raise ValueError(f"Expected one serialized record, but the serializer returned {len(store_models)}.") + return store_models[0] + + def _serialize_record_to_dict(self, record: ModelT) -> dict[str, Any]: + if self.record_type is dict: + source = self._to_builtin_mapping(record) + else: + if self._model_registration is None: + raise RuntimeError(f"Vector model {self.record_type.__name__!r} is not registered.") + source = self._to_builtin_mapping(self._model_registration.encoder(record)) + return self._serialize_mapping_to_store(source) + + @staticmethod + def _to_builtin_mapping(record: Any) -> Mapping[str, Any]: + converted = msgspec.to_builtins(record, str_keys=True, enc_hook=_msgspec_enc_hook) + if not isinstance(converted, Mapping): + raise TypeError("Vector records must serialize to mappings.") + return cast(Mapping[str, Any], converted) + + def _serialize_mapping_to_store(self, source: Mapping[str, Any]) -> dict[str, Any]: + serialized: dict[str, Any] = {} + for field in self.definition.fields: + if field.name in source: + value = source[field.name] + elif field.storage_name is not None and field.storage_name in source: + value = source[field.storage_name] + else: + raise ValueError(f"Record is missing vector store field '{field.name}'.") + serialized[field.storage_name or field.name] = value + return serialized + + async def _add_vectors_to_records(self, records: Sequence[dict[str, Any]]) -> None: + field_generators: list[tuple[VectorStoreField, EmbeddingClient]] = [] + for field in self.definition.vector_fields: + embedding_generator = field.embedding_generator or self.embedding_generator + if embedding_generator is None: + raise ValueError( + f"Vector field '{field.name}' has no embedding generator. " + "Set generate_vectors=False to preserve supplied vector values." + ) + field_generators.append((field, embedding_generator)) + + for field, embedding_generator in field_generators: + storage_name = field.storage_name or field.name + values = [record.get(storage_name) for record in records] + if any(value is None for value in values): + raise ValueError( + f"Vector field '{field.name}' cannot be embedded because at least one value is missing." + ) + options: EmbeddingGenerationOptions = {} + if field.dimensions is not None: + options["dimensions"] = field.dimensions + embeddings = await embedding_generator.get_embeddings(values, options=options) + if len(embeddings) != len(records): + raise IntegrationInvalidResponseException( + f"Embedding client returned {len(embeddings)} vectors for {len(records)} records." + ) + for record, embedding in zip(records, embeddings, strict=True): + record[storage_name] = _normalize_vector(embedding.vector) + + def deserialize( + self, + records: Any | Sequence[Any], + *, + include_vectors: bool = True, + context: Mapping[str, Any] | None = None, + ) -> ModelT | Sequence[ModelT] | None: + """Deserialize one or more backing store records. + + Raises: + TypeError: If a store record has an unsupported type. + ValueError: If records cannot be reconstructed into the requested model shape. + """ + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + if records is None: + return None + is_batch = _is_non_string_sequence(records) + input_records = list(records) if is_batch else [records] + dict_records = self._deserialize_store_models_to_dicts(input_records, context=context) + if not dict_records: + return [] if is_batch else None + deserialized = [ + self._deserialize_dict_to_record(record, include_vectors=include_vectors) for record in dict_records + ] + return deserialized if is_batch else deserialized[0] + + def _deserialize_dict_to_record( + self, + record: Mapping[str, Any], + *, + include_vectors: bool, + ) -> ModelT: + logical_record = self._deserialize_dict_to_mapping(record, include_vectors=include_vectors) + if self.record_type is dict: + return cast(ModelT, logical_record) + if self._model_registration is None: + raise RuntimeError(f"Vector model {self.record_type.__name__!r} is not registered.") + return cast(ModelT, self._model_registration.decoder(logical_record)) + + def _deserialize_dict_to_mapping( + self, + record: Mapping[str, Any], + *, + include_vectors: bool, + ) -> dict[str, Any]: + logical_record: dict[str, Any] = {} + for field in self.definition.fields: + if not include_vectors and field.field_type == "vector": + continue + storage_name = field.storage_name or field.name + if storage_name not in record: + raise IntegrationInvalidResponseException( + f"Vector store response is missing required field '{storage_name}'." + ) + logical_record[field.name] = record[storage_name] + return logical_record + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class BaseVectorCollection(_VectorStoreRecordHandler[KeyT, ModelT], ABC): + """Base class for vector store collection CRUD operations.""" + + def __init__( + self, + record_type: type[ModelT], + *, + definition: VectorStoreCollectionDefinition | None = None, + collection_name: str | None = None, + embedding_generator: EmbeddingClient | None = None, + managed_client: bool = True, + ) -> None: + """Initialize a vector store collection.""" + super().__init__( + record_type, + definition=definition, + embedding_generator=embedding_generator, + ) + self.collection_name = collection_name or self.definition.collection_name or "" + if not self.collection_name: + raise ValueError("A collection name is required when the model definition does not provide one.") + self.managed_client = managed_client + + async def __aenter__(self) -> Self: + """Enter the collection context manager.""" + return self + + async def __aexit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: + """Exit the collection context manager.""" + + @abstractmethod + async def ensure_collection_exists( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Create the collection when it does not exist.""" + ... + + @abstractmethod + async def collection_exists( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> bool: + """Check whether the collection exists.""" + ... + + @abstractmethod + async def ensure_collection_deleted( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Delete the collection when it exists.""" + ... + + @abstractmethod + async def _inner_upsert( + self, + records: Sequence[Any], + *, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[KeyT]: + """Upsert serialized records and return their keys.""" + ... + + @abstractmethod + async def _inner_get( + self, + *, + keys: Sequence[KeyT] | None = None, + top: int = 10, + skip: int = 0, + order_by: Mapping[str, bool] | None = None, + include_vectors: bool = False, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[Any] | None: + """Retrieve store-specific records.""" + ... + + @abstractmethod + async def _inner_delete( + self, + keys: Sequence[KeyT], + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Delete records by key.""" + ... + + async def upsert( + self, + records: Sequence[ModelT], + *, + generate_vectors: bool = True, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[KeyT]: + """Upsert a batch of records. + + Args: + records: A sequence of models. + generate_vectors: Whether to generate vector values, overwriting any supplied values. When ``False``, + supplied values are preserved. + operation_options: Store-specific operation options. + + Returns: + The keys of all upserted records. + + Raises: + TypeError: If record serialization encounters an unsupported type. + ValueError: If record data or returned keys have an invalid shape, or a vector field has no generator. + IntegrationException: If the backing store operation fails. + IntegrationInvalidResponseException: If the backing store returns an unexpected key count. + """ + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + if not _is_non_string_sequence(records): + raise TypeError("records must be a sequence.") + try: + serialized = await self.serialize(records, generate_vectors=generate_vectors) + store_records = list(serialized) if _is_non_string_sequence(serialized) else [serialized] + keys = list(await self._inner_upsert(store_records, operation_options=operation_options)) + except (TypeError, ValueError): + raise + except IntegrationException: + raise + except Exception as exc: + raise IntegrationException( + f"Error upserting records into collection '{self.collection_name}': {exc}" + ) from exc + if len(keys) != len(store_records): + raise IntegrationInvalidResponseException( + f"Expected {len(store_records)} upserted keys, but the store returned {len(keys)}." + ) + return keys + + async def get( + self, + keys: Sequence[KeyT] | None = None, + *, + top: int = 10, + skip: int = 0, + order_by: Mapping[str, bool] | None = None, + include_vectors: bool = False, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[ModelT]: + """Get records by keys or list a page of records. + + Args: + keys: A sequence of keys, or ``None`` to list a page of records. + top: The maximum number of records returned when listing. + skip: The number of records skipped when listing. + order_by: Field names mapped to ascending (``True``) or descending (``False``) order. + include_vectors: Whether returned records include vector fields. + operation_options: Store-specific operation options. + + Returns: + A sequence of models. Keys that do not exist are omitted. + + Raises: + ValueError: If paging arguments are invalid. + TypeError: If keys or a returned record has an unsupported type. + IntegrationException: If retrieval fails. + """ + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + _validate_paging(top=top, skip=skip) + if keys is not None and not _is_non_string_sequence(keys): + raise TypeError("keys must be a sequence.") + try: + records = await self._inner_get( + keys=keys, + top=top, + skip=skip, + order_by=order_by, + include_vectors=include_vectors, + operation_options=operation_options, + ) + except IntegrationException: + raise + except Exception as exc: + raise IntegrationException( + f"Error getting records from collection '{self.collection_name}': {exc}" + ) from exc + if not records: + return [] + deserialized = self.deserialize(records, include_vectors=include_vectors) + return [] if deserialized is None else cast(Sequence[ModelT], deserialized) + + async def delete( + self, + keys: Sequence[KeyT], + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Delete a batch of records by key. + + Args: + keys: The keys to delete. + operation_options: Store-specific operation options. + + Raises: + TypeError: If keys is not a sequence. + IntegrationException: If the backing store operation fails. + """ + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + if not _is_non_string_sequence(keys): + raise TypeError("keys must be a sequence.") + try: + await self._inner_delete(keys, operation_options=operation_options) + except IntegrationException: + raise + except Exception as exc: + raise IntegrationException( + f"Error deleting records from collection '{self.collection_name}': {exc}" + ) from exc + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class BaseVectorStore(ABC): + """Base class for vector stores that create collection clients.""" + + def __init__( + self, + *, + embedding_generator: EmbeddingClient | None = None, + managed_client: bool = True, + ) -> None: + """Initialize a vector store.""" + self.embedding_generator = embedding_generator + self.managed_client = managed_client + + async def __aenter__(self) -> Self: + """Enter the vector store context manager.""" + return self + + async def __aexit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: + """Exit the vector store context manager.""" + + @abstractmethod + def get_collection( + self, + record_type: type[ModelT], + *, + definition: VectorStoreCollectionDefinition | None = None, + collection_name: str | None = None, + embedding_generator: EmbeddingClient | None = None, + ) -> BaseVectorCollection[Any, ModelT]: + """Create a collection client tied to this store.""" + ... + + @abstractmethod + async def list_collection_names( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[str]: + """List collection names.""" + ... + + async def collection_exists( + self, + collection_name: str, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> bool: + """Check whether a collection exists.""" + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + return collection_name in await self.list_collection_names(operation_options=operation_options) + + async def ensure_collection_deleted( + self, + collection_name: str, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Delete a collection when it exists.""" + if not await self.collection_exists(collection_name, operation_options=operation_options): + return + await self._inner_ensure_collection_deleted( + collection_name, + operation_options=operation_options, + ) + + @abstractmethod + async def _inner_ensure_collection_deleted( + self, + collection_name: str, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Delete a collection by name.""" + ... + + +class _LambdaVisitor(NodeVisitor, Generic[FilterT]): + def __init__(self, lambda_parser: Callable[[expr], FilterT]) -> None: + self.lambda_parser = lambda_parser + self.output_filters: list[FilterT] = [] + + def visit_Lambda(self, node: Lambda) -> None: + self.output_filters.append(self.lambda_parser(node.body)) + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class BaseVectorSearch(_VectorStoreRecordHandler[KeyT, ModelT], ABC): + """Base class for vector and keyword-hybrid search.""" + + supported_search_types: ClassVar[set[SearchType]] = {"vector"} + + @abstractmethod + async def _inner_search( + self, + *, + search_type: SearchType, + filter: Any | list[Any] | None = None, + values: Any | None = None, + vector: Vector | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[Any]: + """Execute a search and return raw connector results.""" + ... + + @abstractmethod + def _get_record_from_result(self, result: Any) -> Any: + """Extract a store record from one raw search result.""" + ... + + @abstractmethod + def _get_score_from_result(self, result: Any) -> float | None: + """Extract a score from one raw search result.""" + ... + + @abstractmethod + def _lambda_parser(self, node: AST) -> Any: + """Translate one lambda expression body into a store filter.""" + ... + + @overload + async def search( + self, + values: Any, + *, + search_type: SearchType = "vector", + vector: Vector | None = None, + filter: RecordFilters | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[SearchResponse[ModelT]]: + """Search from a value, optionally with a precomputed vector. + + Args: + values: The value to search for or vectorize. + search_type: Whether to perform vector or keyword-hybrid search. + vector: An optional precomputed query vector. + filter: One or more lambda filters. + top: The maximum number of results. + skip: The number of results to skip. + include_vectors: Whether returned records include vector fields. + vector_property_name: The vector field used for search. + additional_property_name: The data field used for keyword-hybrid search. + score_threshold: The minimum similarity or maximum distance accepted. + Results without scores remain included. + operation_options: Store-specific operation options. + + Returns: + Lazily consumed search results. + + Raises: + ValueError: If paging or search arguments are invalid. + NotImplementedError: If the search type is unsupported. + IntegrationException: If vector generation or search fails. + """ + ... + + @overload + async def search( + self, + *, + search_type: Literal["vector"] = "vector", + vector: Vector, + filter: RecordFilters | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[SearchResponse[ModelT]]: + """Search from a required precomputed vector. + + Args: + search_type: The vector search type. + vector: The precomputed query vector. + filter: One or more lambda filters. + top: The maximum number of results. + skip: The number of results to skip. + include_vectors: Whether returned records include vector fields. + vector_property_name: The vector field used for search. + additional_property_name: The data field used for keyword-hybrid search. + score_threshold: The minimum similarity or maximum distance accepted. + Results without scores remain included. + operation_options: Store-specific operation options. + + Returns: + Lazily consumed search results. + + Raises: + ValueError: If paging or search arguments are invalid. + NotImplementedError: If vector search is unsupported. + IntegrationException: If search execution fails. + """ + ... + + async def search( + self, + values: Any | None = None, + *, + search_type: SearchType = "vector", + vector: Vector | None = None, + filter: RecordFilters | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[SearchResponse[ModelT]]: + """Search the vector store. + + Args: + values: The value to search for or vectorize. + search_type: Whether to perform vector or keyword-hybrid search. + vector: A precomputed query vector. + filter: One or more lambda filters. + top: The maximum number of results. + skip: The number of results to skip. + include_vectors: Whether returned records include vector fields. + vector_property_name: The vector field used for search. + additional_property_name: The data field used for keyword-hybrid search. + score_threshold: The minimum similarity or maximum distance accepted. + Results without scores remain included. + operation_options: Store-specific operation options. + + Returns: + Lazily consumed search results. + + Raises: + ValueError: If paging or search arguments are invalid. + NotImplementedError: If the search type is unsupported. + IntegrationException: If the backing store search fails. + """ + mark_feature_used(FeatureIndex.CORE_VECTOR_STORES) + if search_type not in ("vector", "keyword_hybrid"): + raise ValueError(f"Unknown search type '{search_type}'.") + if search_type not in self.supported_search_types: + raise NotImplementedError(f"Search type '{search_type}' is not supported by {type(self).__name__}.") + if values is None and vector is None: + raise ValueError("Search requires values or a precomputed vector.") + if search_type == "keyword_hybrid" and values is None: + raise ValueError("Keyword-hybrid search requires values.") + + _validate_paging(top=top, skip=skip) + try: + self._validate_score_threshold( + score_threshold=score_threshold, + vector_property_name=vector_property_name, + ) + resolved_vector = vector + if resolved_vector is None and values is not None: + resolved_vector = await self._generate_vector_from_values( + values, + vector_property_name=vector_property_name, + ) + translated_filter = self._build_filter(filter) + raw_results = await self._inner_search( + search_type=search_type, + filter=translated_filter, + values=values, + vector=resolved_vector, + top=top, + skip=skip, + include_vectors=include_vectors, + vector_property_name=vector_property_name, + additional_property_name=additional_property_name, + score_threshold=score_threshold, + operation_options=operation_options, + ) + return SearchResults( + self._get_search_results_from_results( + raw_results.results, + include_vectors=include_vectors, + vector_property_name=vector_property_name, + score_threshold=score_threshold, + ), + metadata=raw_results.metadata, + ) + except (TypeError, ValueError): + raise + except IntegrationException: + raise + except Exception as exc: + raise IntegrationException(f"Vector search failed: {exc}") from exc + + def _validate_score_threshold( + self, + *, + score_threshold: float | None, + vector_property_name: str | None, + ) -> None: + if score_threshold is None: + return + vector_field = self.definition.try_get_vector_field(vector_property_name) + if vector_field is None: + raise ValueError("A score threshold requires a vector field.") + if vector_field.distance_function == "DEFAULT": + raise ValueError("A score threshold requires an explicit distance function on the vector field.") + + async def _generate_vector_from_values( + self, + values: Any, + *, + vector_property_name: str | None, + ) -> Vector | None: + vector_field = self.definition.try_get_vector_field(vector_property_name) + if vector_field is None: + if vector_property_name is not None: + raise ValueError(f"Vector field '{vector_property_name}' was not found in the collection definition.") + return None + embedding_generator = vector_field.embedding_generator or self.embedding_generator + if embedding_generator is None: + return None + embedding_options: EmbeddingGenerationOptions = {} + if vector_field.dimensions is not None: + embedding_options["dimensions"] = vector_field.dimensions + embeddings = await embedding_generator.get_embeddings([values], options=embedding_options) + if len(embeddings) != 1: + raise IntegrationInvalidResponseException( + f"Embedding client returned {len(embeddings)} vectors for one search value." + ) + generated_vector = embeddings[0].vector + return _normalize_vector(generated_vector) + + def _build_filter(self, search_filter: RecordFilters | None) -> Any | list[Any] | None: + """Translate lambda filters with the connector's AST parser.""" + if not search_filter: + return None + filters: list[RecordFilter] + if _is_non_string_sequence(search_filter) and not callable(search_filter): + filters = cast(list[RecordFilter], list(search_filter)) + else: + filters = [cast(RecordFilter, search_filter)] + visitor = _LambdaVisitor(self._lambda_parser) + try: + for filter_item in filters: + source = ( + filter_item + if isinstance(filter_item, str) + else getsource(cast(Callable[..., Any], filter_item)).strip() + ) + visitor.visit(parse(source)) + except (OSError, SyntaxError, TypeError) as exc: + raise ValueError(f"Unable to parse vector search filter: {exc}") from exc + if not visitor.output_filters: + raise ValueError("No lambda expression was found in the vector search filter.") + return visitor.output_filters[0] if len(visitor.output_filters) == 1 else visitor.output_filters + + def _get_search_results_from_results( + self, + results: AsyncIterable[Any] | Sequence[Any], + *, + include_vectors: bool, + vector_property_name: str | None, + score_threshold: float | None, + ) -> AsyncIterable[SearchResponse[ModelT]]: + """Convert raw connector results into deserialized search responses.""" + + async def generate() -> AsyncIterator[SearchResponse[ModelT]]: + try: + async for result in _as_async_iterable(results): + try: + record = self.deserialize( + self._get_record_from_result(result), + include_vectors=include_vectors, + ) + if record is None or _is_non_string_sequence(record): + if record is None: + continue + raise IntegrationInvalidResponseException( + "A search result must deserialize to exactly one record." + ) + score = self._get_score_from_result(result) + if not self._meets_score_threshold( + score, + score_threshold=score_threshold, + vector_property_name=vector_property_name, + ): + continue + yield SearchResponse(record=cast(ModelT, record), score=score) + except IntegrationException: + raise + except Exception as exc: + raise IntegrationInvalidResponseException( + f"Vector search result conversion failed: {exc}" + ) from exc + except IntegrationException: + raise + except Exception as exc: + raise IntegrationException(f"Vector search iteration failed: {exc}") from exc + + return generate() + + def _meets_score_threshold( + self, + score: float | None, + *, + score_threshold: float | None, + vector_property_name: str | None, + ) -> bool: + """Apply a threshold when a result includes a comparable score. + + Results without scores remain included because the threshold cannot be + evaluated for them. + """ + if score_threshold is None or score is None: + return True + vector_field = self.definition.try_get_vector_field(vector_property_name) + if vector_field is None or vector_field.distance_function is None: + return True + comparison = DISTANCE_FUNCTION_DIRECTION_HELPER.get(vector_field.distance_function) + return comparison(score, score_threshold) if comparison is not None else True + + +@runtime_checkable +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class SupportsVectorUpsert(Protocol[KeyT, ModelT]): + """Protocol for vector collection CRUD operations.""" + + collection_name: str + record_type: type[ModelT] + definition: VectorStoreCollectionDefinition + + async def upsert( + self, + records: Sequence[ModelT], + *, + generate_vectors: bool = True, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[KeyT]: + """Upsert a batch of records, generating embeddings by default.""" + ... + + async def get( + self, + keys: Sequence[KeyT] | None = None, + *, + top: int = 10, + skip: int = 0, + order_by: Mapping[str, bool] | None = None, + include_vectors: bool = False, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[ModelT]: + """Get records by keys or list a page of records, excluding vectors by default.""" + ... + + async def delete( + self, + keys: Sequence[KeyT], + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + """Delete a batch of records by key.""" + ... + + +@runtime_checkable +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +class SupportsVectorSearch(Protocol[ModelT]): + """Protocol for vector and keyword-hybrid search.""" + + @overload + async def search( + self, + values: Any, + *, + search_type: SearchType = "vector", + vector: Vector | None = None, + filter: RecordFilters | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[SearchResponse[ModelT]]: + """Search from a value, optionally with a precomputed vector. + + Args: + values: The value to search for or vectorize. + search_type: Whether to perform vector or keyword-hybrid search. + vector: An optional precomputed query vector. + filter: One or more lambda filters. + top: The maximum number of results. + skip: The number of results to skip. + include_vectors: Whether returned records include vector fields. + vector_property_name: The vector field used for search. + additional_property_name: The data field used for keyword-hybrid search. + score_threshold: The minimum similarity or maximum distance accepted. + Results without scores remain included. + operation_options: Store-specific operation options. + + Returns: + Lazily consumed search results. + + Raises: + ValueError: If paging or search arguments are invalid. + NotImplementedError: If the search type is unsupported. + IntegrationException: If vector generation or search fails. + """ + ... + + @overload + async def search( + self, + *, + search_type: Literal["vector"] = "vector", + vector: Vector, + filter: RecordFilters | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[SearchResponse[ModelT]]: + """Search from a required precomputed vector. + + Args: + search_type: The vector search type. + vector: The precomputed query vector. + filter: One or more lambda filters. + top: The maximum number of results. + skip: The number of results to skip. + include_vectors: Whether returned records include vector fields. + vector_property_name: The vector field used for search. + additional_property_name: The data field used for keyword-hybrid search. + score_threshold: The minimum similarity or maximum distance accepted. + Results without scores remain included. + operation_options: Store-specific operation options. + + Returns: + Lazily consumed search results. + + Raises: + ValueError: If paging or search arguments are invalid. + NotImplementedError: If vector search is unsupported. + IntegrationException: If search execution fails. + """ + ... + + +@experimental(feature_id=ExperimentalFeature.VECTOR_STORES) +def create_vector_search_tool( + search: SupportsVectorSearch[ModelT], + *, + name: str = _DEFAULT_SEARCH_TOOL_NAME, + description: str = _DEFAULT_SEARCH_TOOL_DESCRIPTION, + approval_mode: Literal["always_require", "never_require"] = "never_require", + search_type: SearchType = "vector", + parameters: type[BaseModel] | Mapping[str, Any] | None = None, + top: int = 5, + skip: int = 0, + filter: RecordFilters | None = None, + filter_mapper: Callable[[RecordFilters | None, Mapping[str, Any]], RecordFilters | None] | None = None, + result_mapper: Callable[[SearchResponse[ModelT]], str | Content | Sequence[Content]] | None = None, +) -> FunctionTool: + """Create an agent-usable tool backed by vector search. + + Args: + search: The vector search capability invoked by the tool. + name: The tool name. + description: The tool description shown to the model. + approval_mode: Whether the tool requires approval before invocation. + search_type: Whether the tool performs vector or keyword-hybrid search. + parameters: A Pydantic model or JSON schema declaring the tool parameters. + It must declare ``query`` as a required string. A custom schema can + expose ``top`` and ``skip`` as integers with finite ``maximum`` values; + additional fields are passed to ``filter_mapper``. + top: The default result limit and the maximum when ``parameters`` does not expose ``top``. + skip: The default offset and the maximum when ``parameters`` does not expose ``skip``. + filter: A fixed filter applied to each tool invocation. + filter_mapper: Maps additional declared tool arguments to search filters. + The default creates equality filters for each additional argument. + result_mapper: Maps each search response to text or one or more multimodal content items. + + Returns: + A function tool with only a ``query`` parameter by default. Custom parameters can expose + ``top``, ``skip``, and fields mapped into filters by ``filter_mapper``. + + Raises: + ValueError: If parameters or paging limits are invalid. + NotImplementedError: If the search type is unsupported. + """ + _validate_paging(top=top, skip=skip) + map_filter = filter_mapper or _default_search_filter_mapper + map_result = result_mapper or _default_search_result_mapper + input_model = parameters if parameters is not None else _default_search_tool_parameters() + max_top, max_skip = _validate_search_tool_parameters( + input_model, + default_top=top, + default_skip=skip, + ) + + async def search_tool(**arguments: Any) -> list[Content]: + query = arguments.pop("query") + if not isinstance(query, str): + raise TypeError("The search tool 'query' argument must be a string.") + invocation_top = arguments.pop("top", top) + invocation_skip = arguments.pop("skip", skip) + _validate_paging(top=invocation_top, skip=invocation_skip) + if invocation_top > max_top: + raise ValueError(f"top must not exceed the configured maximum of {max_top}.") + if invocation_skip > max_skip: + raise ValueError(f"skip must not exceed the configured maximum of {max_skip}.") + dynamic_filter = map_filter(filter, arguments) + results = await search.search( + query, + search_type=search_type, + filter=dynamic_filter, + top=invocation_top, + skip=invocation_skip, + ) + mapped_results: list[Content] = [] + consumed_results = 0 + async for result in results: + if consumed_results >= invocation_top: + break + consumed_results += 1 + mapped = map_result(result) + if isinstance(mapped, str): + mapped_results.append(Content.from_text(mapped)) + elif isinstance(mapped, Content): + mapped_results.append(mapped) + else: + mapped_results.extend(mapped) + return mapped_results + + return FunctionTool( + name=name, + description=description, + approval_mode=approval_mode, + func=search_tool, + input_model=input_model, + ) + + +def _is_non_string_sequence(value: Any) -> TypeGuard[Sequence[Any]]: + return isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray, Mapping)) + + +async def _as_async_iterable( + values: AsyncIterable[ResultT] | Sequence[ResultT], +) -> AsyncIterator[ResultT]: + if isinstance(values, AsyncIterable): + async for value in values: + yield value + return + for value in values: + yield value + + +def _default_search_tool_parameters() -> dict[str, Any]: + return { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "The query to search for.", + }, + }, + "required": ["query"], + "additionalProperties": False, + } + + +def _validate_search_tool_parameters( + parameters: type[BaseModel] | Mapping[str, Any], + *, + default_top: int, + default_skip: int, +) -> tuple[int, int]: + schema: Mapping[str, Any] = parameters.model_json_schema() if isinstance(parameters, type) else parameters + raw_properties = schema.get("properties") + if not isinstance(raw_properties, Mapping): + raise ValueError("Search tool parameters must define object properties.") + properties = cast(Mapping[str, Any], raw_properties) + query_schema = properties.get("query") + required = schema.get("required") + query_type = cast(Mapping[str, Any], query_schema).get("type") if isinstance(query_schema, Mapping) else None + if ( + not isinstance(query_schema, Mapping) + or query_type != "string" + or not _is_non_string_sequence(required) + or "query" not in required + ): + raise ValueError("Search tool parameters must define 'query' as a required string.") + + limits = {"top": default_top, "skip": default_skip} + for name, minimum in (("top", 1), ("skip", 0)): + parameter_schema = properties.get(name) + if parameter_schema is None: + continue + if not isinstance(parameter_schema, Mapping): + raise ValueError(f"Search tool parameter '{name}' must be an integer.") + typed_parameter_schema = cast(Mapping[str, Any], parameter_schema) + if typed_parameter_schema.get("type") != "integer": + raise ValueError(f"Search tool parameter '{name}' must be an integer.") + maximum = typed_parameter_schema.get("maximum") + if not isinstance(maximum, int) or isinstance(maximum, bool) or maximum < minimum: + raise ValueError(f"Search tool parameter '{name}' must declare an integer maximum of at least {minimum}.") + configured_default = default_top if name == "top" else default_skip + if configured_default > maximum: + raise ValueError(f"Configured {name}={configured_default} exceeds the parameter maximum of {maximum}.") + limits[name] = maximum + return limits["top"], limits["skip"] + + +def _default_search_filter_mapper( + search_filter: RecordFilters | None, + arguments: Mapping[str, Any], +) -> RecordFilters | None: + dynamic_filters: list[RecordFilter] = [] + for name, value in arguments.items(): + if not name.isidentifier(): + raise ValueError(f"Search tool parameter '{name}' cannot be mapped to a model field.") + dynamic_filters.append(f"lambda record: record.{name} == {value!r}") + if not dynamic_filters: + return search_filter + if search_filter is None: + return dynamic_filters + if _is_non_string_sequence(search_filter) and not callable(search_filter): + return [*cast(Sequence[RecordFilter], search_filter), *dynamic_filters] + return [cast(RecordFilter, search_filter), *dynamic_filters] + + +def _default_search_result_mapper(response: SearchResponse[Any]) -> str: + return msgspec.json.encode( + response, + enc_hook=_msgspec_enc_hook, + ).decode() diff --git a/python/packages/core/tests/core/test_vectors.py b/python/packages/core/tests/core/test_vectors.py new file mode 100644 index 00000000000..9bb045e61d7 --- /dev/null +++ b/python/packages/core/tests/core/test_vectors.py @@ -0,0 +1,1373 @@ +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +import warnings +from ast import AST, unparse +from collections.abc import AsyncIterable, Mapping, Sequence +from dataclasses import FrozenInstanceError, dataclass, field +from typing import Annotated, Any, ClassVar, cast +from unittest.mock import patch + +import msgspec +import pytest +from pydantic import BaseModel +from pydantic import Field as PydanticField +from typing_extensions import TypeVar + +from agent_framework import ( + DISTANCE_FUNCTION_DIRECTION_HELPER, + BaseEmbeddingClient, + BaseVectorCollection, + BaseVectorSearch, + BaseVectorStore, + Content, + DistanceFunction, + Embedding, + EmbeddingGenerationOptions, + ExperimentalFeature, + FieldTypes, + GeneratedEmbeddings, + IndexKind, + SearchResponse, + SearchResults, + SearchType, + SupportsVectorSearch, + SupportsVectorUpsert, + VectorStoreCollectionDefinition, + VectorStoreField, + create_vector_search_tool, + register_vectorstoremodel, + vectorstoremodel, +) +from agent_framework._feature_stage import ExperimentalWarning +from agent_framework._telemetry import FeatureIndex +from agent_framework._vectors import _VectorStoreRecordHandler as VectorStoreRecordHandler +from agent_framework.exceptions import IntegrationException, IntegrationInvalidResponseException + +pytestmark = pytest.mark.filterwarnings("ignore::agent_framework._feature_stage.ExperimentalWarning") + +with warnings.catch_warnings(): + warnings.simplefilter("ignore", ExperimentalWarning) + RecordVector = Annotated[ + str | list[float] | None, + VectorStoreField( + "vector", + dimensions=2, + index_kind="hnsw", + distance_function="cosine_similarity", + ), + ] + + @vectorstoremodel(collection_name="records") + @dataclass + class Record: + id: Annotated[str, VectorStoreField("key", storage_name="record_id")] + text: Annotated[str, VectorStoreField("data", storage_name="body", is_full_text_indexed=True)] + vector: RecordVector = None + category: str = "general" + + +class MockEmbeddingClient(BaseEmbeddingClient): + def __init__(self) -> None: + super().__init__() + self.values: list[Any] = [] + self.options: EmbeddingGenerationOptions | None = None + + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[list[float]]: + self.values = list(values) + self.options = options + return GeneratedEmbeddings([Embedding(vector=[float(len(str(value))), 0.5]) for value in values]) + + +class MockCollection(BaseVectorCollection[str, Record], BaseVectorSearch[str, Record]): + supported_key_types: ClassVar[set[str] | None] = {"str"} + supported_vector_types: ClassVar[set[str] | None] = {"float"} + supported_search_types: ClassVar[set[SearchType]] = {"vector", "keyword_hybrid"} + + def __init__(self, *, embedding_generator: MockEmbeddingClient | None = None) -> None: + super().__init__(Record, embedding_generator=embedding_generator) + self.created = False + self.records: dict[str, dict[str, Any]] = {} + self.last_search_type: str | None = None + self.last_search_vector: Sequence[float | int] | None = None + self.last_search_filter: Any | list[Any] | None = None + self.last_search_top = 0 + self.last_search_skip = 0 + self.fail_upsert = False + self.upsert_error: Exception | None = None + self.get_error: Exception | None = None + self.delete_error: Exception | None = None + self.search_error: Exception | None = None + self.upsert_keys: Sequence[str] | None = None + self.raw_search_results: AsyncIterable[Any] | Sequence[Any] | None = None + + async def ensure_collection_exists( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + self.created = True + + async def collection_exists( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> bool: + return self.created + + async def ensure_collection_deleted( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + self.created = False + self.records.clear() + + async def _inner_upsert( + self, + records: Sequence[Any], + *, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[str]: + if self.upsert_error is not None: + raise self.upsert_error + if self.fail_upsert: + raise RuntimeError("store unavailable") + keys: list[str] = [] + for record in records: + mapping = cast(Mapping[str, Any], record) + key = cast(str, mapping["record_id"]) + self.records[key] = dict(mapping) + keys.append(key) + return self.upsert_keys if self.upsert_keys is not None else keys + + async def _inner_get( + self, + *, + keys: Sequence[str] | None = None, + top: int = 10, + skip: int = 0, + order_by: Mapping[str, bool] | None = None, + include_vectors: bool = False, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[Any] | None: + if self.get_error is not None: + raise self.get_error + if keys is not None: + return [self.records[key] for key in keys if key in self.records] + return list(self.records.values())[skip : skip + top] + + async def _inner_delete( + self, + keys: Sequence[str], + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + if self.delete_error is not None: + raise self.delete_error + for key in keys: + self.records.pop(key, None) + + async def _inner_search( + self, + *, + search_type: SearchType, + filter: Any | list[Any] | None = None, + values: Any | None = None, + vector: Sequence[float | int] | None = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[Any]: + if self.search_error is not None: + raise self.search_error + self.last_search_type = search_type + self.last_search_vector = vector + self.last_search_filter = filter + self.last_search_top = top + self.last_search_skip = skip + raw_results = self.raw_search_results or [ + {"record": record, "score": score} for record, score in zip(self.records.values(), (0.9, 0.4), strict=False) + ] + return SearchResults(raw_results, metadata={"mock_count": len(self.records)}) + + def _get_record_from_result(self, result: Any) -> Any: + return result["record"] + + def _get_score_from_result(self, result: Any) -> float | None: + return cast(float | None, result["score"]) + + def _lambda_parser(self, node: AST) -> str: + return unparse(node) + + +StoreModelT = TypeVar("StoreModelT") + + +class MockStore(BaseVectorStore): + def __init__(self, collection: MockCollection) -> None: + super().__init__() + self.collection = collection + + def get_collection( + self, + record_type: type[StoreModelT], + *, + definition: VectorStoreCollectionDefinition | None = None, + collection_name: str | None = None, + embedding_generator: Any | None = None, + ) -> BaseVectorCollection[Any, StoreModelT]: + return cast(BaseVectorCollection[Any, StoreModelT], self.collection) + + async def list_collection_names( + self, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> Sequence[str]: + return [self.collection.collection_name] if self.collection.created else [] + + async def _inner_ensure_collection_deleted( + self, + collection_name: str, + *, + operation_options: Mapping[str, Any] | None = None, + ) -> None: + assert collection_name == self.collection.collection_name + await self.collection.ensure_collection_deleted(operation_options=operation_options) + + +def test_vector_literal_types_and_distance_directions() -> None: + field_type: FieldTypes = "vector" + index_kind: IndexKind = "hnsw" + distance_function: DistanceFunction = "cosine_similarity" + + assert field_type == "vector" + assert index_kind == "hnsw" + assert distance_function == "cosine_similarity" + assert DISTANCE_FUNCTION_DIRECTION_HELPER["cosine_similarity"](0.5, 0.5) + assert DISTANCE_FUNCTION_DIRECTION_HELPER["cosine_distance"](0.5, 0.5) + assert not DISTANCE_FUNCTION_DIRECTION_HELPER["cosine_distance"](0.6, 0.5) + + +def test_vector_apis_are_marked_experimental() -> None: + staged_apis = ( + VectorStoreField, + VectorStoreCollectionDefinition, + vectorstoremodel, + SearchResponse, + SearchResults, + BaseVectorCollection, + BaseVectorStore, + BaseVectorSearch, + register_vectorstoremodel, + ) + for api in staged_apis: + assert getattr(api, "__feature_stage__", None) == "experimental" + assert getattr(api, "__feature_id__", None) == ExperimentalFeature.VECTOR_STORES.value + assert ".. warning:: Experimental" in (api.__doc__ or "") + + staged_protocols = ( + SupportsVectorUpsert, + SupportsVectorSearch, + ) + for protocol in staged_protocols: + assert ".. warning:: Experimental" in (protocol.__doc__ or "") + + +def test_vector_field_validates_vector_options() -> None: + with pytest.raises(ValueError, match="positive"): + cast(Any, VectorStoreField)("vector") + with pytest.raises(ValueError, match="Vector-only"): + cast(Any, VectorStoreField)("data", dimensions=3) + with pytest.raises(ValueError, match="index kind"): + cast(Any, VectorStoreField)("vector", dimensions=3, index_kind="unknown") + with pytest.raises(ValueError, match="distance function"): + cast(Any, VectorStoreField)("vector", dimensions=3, distance_function="unknown") + + +def test_collection_definition_exposes_fields() -> None: + definition = cast(VectorStoreCollectionDefinition, vars(Record)["__vectorstoremodel_definition__"]) + + assert definition.collection_name == "records" + assert definition.key_name == "id" + assert definition.key_field_storage_name == "record_id" + assert definition.names == ["id", "text", "vector"] + assert definition.storage_names == ["record_id", "body", "vector"] + assert definition.data_field_names == ["text"] + assert definition.vector_field_names == ["vector"] + assert definition.get_names(include_vector_fields=False) == ["id", "text"] + assert definition.get_storage_names(include_key_field=False) == ["body", "vector"] + assert isinstance(definition.fields, tuple) + assert definition.vector_fields[0].dimensions == 2 + assert definition.vector_fields[0].index_kind == "hnsw" + assert definition.vector_fields[0].distance_function == "cosine_similarity" + + frozen_field = cast(Any, definition.fields[0]) + with pytest.raises(FrozenInstanceError): + frozen_field.name = "changed" + frozen_definition = cast(Any, definition) + with pytest.raises(FrozenInstanceError): + frozen_definition.fields = () + + +@pytest.mark.parametrize( + "fields, message", + [ + ([], "at least one"), + ([VectorStoreField("data", name="text")], "exactly one key"), + ( + [ + VectorStoreField("key", name="id"), + VectorStoreField("key", name="other_id"), + ], + "exactly one key", + ), + ( + [ + VectorStoreField("key", name="id"), + VectorStoreField("data", name="id"), + ], + "must be unique", + ), + ], +) +def test_collection_definition_rejects_invalid_fields( + fields: list[VectorStoreField], + message: str, +) -> None: + with pytest.raises(ValueError, match=message): + VectorStoreCollectionDefinition(fields) + + +def test_vectorstoremodel_supports_pydantic_models() -> None: + @vectorstoremodel + class PydanticRecord(BaseModel): + id: Annotated[str, VectorStoreField("key")] + vector: Annotated[list[float] | None, VectorStoreField("vector", dimensions=2)] = None + + definition = cast( + VectorStoreCollectionDefinition, + vars(PydanticRecord)["__vectorstoremodel_definition__"], + ) + assert vars(PydanticRecord)["__vectorstoremodel__"] + assert definition.key_field.type_ == "str" + assert definition.vector_fields[0].type_ == "float" + handler = VectorStoreRecordHandler(PydanticRecord) + record = handler.deserialize({"id": "one", "vector": [1.0, 0.0]}, include_vectors=False) + assert isinstance(record, PydanticRecord) + assert record.vector is None + + +def test_vectorstoremodel_supports_plain_classes() -> None: + @vectorstoremodel + class PlainRecord: + def __init__( + self, + id: Annotated[str, VectorStoreField("key")], + text: Annotated[str, VectorStoreField("data")], + ) -> None: + self.id = id + self.text = text + + definition = cast( + VectorStoreCollectionDefinition, + vars(PlainRecord)["__vectorstoremodel_definition__"], + ) + assert definition.names == ["id", "text"] + + +def test_vectorstoremodel_ignores_fields_with_defaults() -> None: + assert ( + "category" + not in cast( + VectorStoreCollectionDefinition, + vars(Record)["__vectorstoremodel_definition__"], + ).names + ) + + +def test_vectorstoremodel_detects_factory_and_required_slotted_defaults() -> None: + @vectorstoremodel + @dataclass(slots=True) + class FactoryRecord: + id: Annotated[str, VectorStoreField("key")] + ignored: list[str] = field(default_factory=list) + + assert ( + "ignored" + not in cast( + VectorStoreCollectionDefinition, + vars(FactoryRecord)["__vectorstoremodel_definition__"], + ).names + ) + + class InvalidStruct(msgspec.Struct): + id: Annotated[str, VectorStoreField("key")] + required_but_unmapped: str + + with pytest.raises(ValueError, match="required_but_unmapped"): + vectorstoremodel(InvalidStruct) + + class RequiredVector(msgspec.Struct): + id: Annotated[str, VectorStoreField("key")] + vector: Annotated[list[float], VectorStoreField("vector", dimensions=2)] + + with pytest.raises(ValueError, match="must declare defaults"): + vectorstoremodel(RequiredVector) + + +def test_vectorstoremodel_rejects_required_unmapped_fields() -> None: + class InvalidRecord: + id: Annotated[str, VectorStoreField("key")] + required_but_unmapped: str + + with pytest.raises(ValueError, match="required_but_unmapped"): + vectorstoremodel(InvalidRecord) + assert not hasattr(InvalidRecord, "__vectorstoremodel__") + + +async def test_collection_and_search_validate_paging() -> None: + with pytest.raises(ValueError, match="greater than zero"): + await MockCollection().get(top=0) + with pytest.raises(ValueError, match="negative"): + await MockCollection().search("query", skip=-1) + + +def test_record_handler_validates_connector_field_types() -> None: + class IntKeyHandler(VectorStoreRecordHandler[str, Record]): + supported_key_types: ClassVar[set[str] | None] = {"int"} + + with pytest.raises(ValueError, match="Key field type"): + IntKeyHandler(Record) + + +async def test_record_handler_serializes_dict_records_with_explicit_definition() -> None: + definition = VectorStoreCollectionDefinition([ + VectorStoreField("key", name="id", storage_name="record_id"), + VectorStoreField("data", name="text", storage_name="body"), + ]) + handler = VectorStoreRecordHandler(dict, definition=definition) + + serialized = await handler.serialize({"id": "one", "text": "hello"}) + assert serialized == {"record_id": "one", "body": "hello"} + assert handler.deserialize(serialized) == {"id": "one", "text": "hello"} + assert handler.deserialize([]) == [] + + with pytest.raises(IntegrationInvalidResponseException, match="missing required field 'body'"): + handler.deserialize({"record_id": "one"}) + assert handler.deserialize({"record_id": "one", "body": None}) == {"id": "one", "text": None} + + with pytest.raises(ValueError, match="missing.*text"): + await handler.serialize({"id": "missing-text"}) + + +async def test_batch_serializer_preserves_cardinality() -> None: + class DroppingHandler(VectorStoreRecordHandler[Any, Record]): + def _serialize_dicts_to_store_models( + self, + records: Sequence[dict[str, Any]], + *, + context: Mapping[str, Any] | None = None, + ) -> Sequence[Any]: + return records[:-1] + + with pytest.raises(IntegrationInvalidResponseException, match="Expected 2 serialized records"): + await DroppingHandler(Record).serialize( + [ + Record("one", "first"), + Record("two", "second"), + ], + generate_vectors=False, + ) + + +async def test_record_handler_supports_msgspec_structs() -> None: + @vectorstoremodel + class MsgspecRecord(msgspec.Struct): + id: Annotated[str, VectorStoreField("key")] + vector: Annotated[list[float] | None, VectorStoreField("vector", dimensions=2)] = None + + handler = VectorStoreRecordHandler(MsgspecRecord) + serialized = await handler.serialize(MsgspecRecord("one", [1.0, 0.0]), generate_vectors=False) + deserialized = handler.deserialize(serialized) + + assert serialized == {"id": "one", "vector": [1.0, 0.0]} + assert deserialized == MsgspecRecord("one", [1.0, 0.0]) + + +async def test_record_handler_uses_registered_codecs() -> None: + @dataclass + class CustomRecord: + id: str + text: str + + definition = VectorStoreCollectionDefinition( + [ + VectorStoreField("key", name="id", storage_name="record_id"), + VectorStoreField("data", name="text", storage_name="body"), + ], + ) + register_vectorstoremodel( + CustomRecord, + definition=definition, + encoder=lambda record: {"id": record.id, "text": record.text.upper()}, + decoder=lambda record: CustomRecord(**record), + ) + handler = VectorStoreRecordHandler(CustomRecord) + + serialized = await handler.serialize(CustomRecord("one", "hello")) + assert serialized == {"record_id": "one", "body": "HELLO"} + assert handler.deserialize(serialized) == CustomRecord("one", "HELLO") + + +async def test_register_vectorstoremodel_supports_independent_encoder_override() -> None: + @dataclass + class RegisteredRecord: + id: str = "" + + definition = VectorStoreCollectionDefinition([VectorStoreField("key", name="id")]) + + def encoder(record: RegisteredRecord) -> Mapping[str, Any]: + return {"id": record.id} + + register_vectorstoremodel(RegisteredRecord, definition=definition, encoder=encoder) + handler = VectorStoreRecordHandler(RegisteredRecord) + assert await handler.serialize(RegisteredRecord("one")) == {"id": "one"} + assert handler.deserialize({"id": "one"}) == RegisteredRecord("one") + + with pytest.raises(ValueError, match="another definition"): + register_vectorstoremodel( + RegisteredRecord, + definition=VectorStoreCollectionDefinition([VectorStoreField("key", name="other_id")]), + ) + + +async def test_array_like_vectors_round_trip_without_array_dependency() -> None: + class ArrayLike: + __slots__ = ("values",) + + def __init__(self, values: list[float]) -> None: + self.values = values + + def tolist(self) -> list[float]: + return self.values + + def decode_array_record(record: Mapping[str, Any]) -> ArrayRecord: + return ArrayRecord( + id=cast(str, record["id"]), + vector=ArrayLike(cast(list[float], record["vector"])), + ) + + @vectorstoremodel(decoder=decode_array_record) + @dataclass + class ArrayRecord: + id: Annotated[str, VectorStoreField("key")] + vector: Annotated[Any, VectorStoreField("vector", dimensions=3)] + + handler = VectorStoreRecordHandler(ArrayRecord) + serialized = await handler.serialize( + ArrayRecord("one", ArrayLike([0.1, 0.2, 0.3])), + generate_vectors=False, + ) + restored = handler.deserialize(serialized) + + assert serialized == {"id": "one", "vector": [0.1, 0.2, 0.3]} + assert isinstance(restored, ArrayRecord) + assert restored.vector.values == [0.1, 0.2, 0.3] + + +async def test_custom_encoder_normalizes_array_like_vectors() -> None: + class ArrayLike: + def tolist(self) -> list[float]: + return [0.1, 0.2, 0.3] + + @dataclass + class CustomArrayRecord: + id: str + vector: ArrayLike + + definition = VectorStoreCollectionDefinition([ + VectorStoreField("key", name="id"), + VectorStoreField("vector", name="vector", dimensions=3), + ]) + register_vectorstoremodel( + CustomArrayRecord, + definition=definition, + encoder=lambda record: {"id": record.id, "vector": record.vector}, + decoder=lambda record: CustomArrayRecord( + id=cast(str, record["id"]), + vector=ArrayLike(), + ), + ) + + serialized = await VectorStoreRecordHandler(CustomArrayRecord).serialize( + CustomArrayRecord("one", ArrayLike()), + generate_vectors=False, + ) + assert serialized == {"id": "one", "vector": [0.1, 0.2, 0.3]} + + +async def test_pydantic_aliases_round_trip_by_field_name() -> None: + @vectorstoremodel + class AliasedRecord(BaseModel): + id: Annotated[str, PydanticField(alias="record_id"), VectorStoreField("key")] + + handler = VectorStoreRecordHandler(AliasedRecord) + serialized = await handler.serialize(AliasedRecord.model_validate({"record_id": "one"})) + restored = handler.deserialize(serialized) + + assert serialized == {"id": "one"} + assert isinstance(restored, AliasedRecord) + assert restored.id == "one" + + +async def test_collection_serializes_records_and_generates_vectors() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + + serialized = await collection.serialize(Record("one", "hello", "embed this")) + + assert serialized == { + "record_id": "one", + "body": "hello", + "vector": [10.0, 0.5], + } + assert embedding_client.values == ["embed this"] + assert embedding_client.options == {"dimensions": 2} + + +async def test_upsert_controls_embedding_generation() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + + await collection.upsert([Record("generated", "text", [1.0, 0.0])]) + + assert embedding_client.values == [[1.0, 0.0]] + assert collection.records["generated"]["vector"] == [10.0, 0.5] + + embedding_client.values.clear() + await collection.upsert( + [Record("preserved", "text", [1.0, 0.0])], + generate_vectors=False, + ) + + assert embedding_client.values == [] + assert collection.records["preserved"]["vector"] == [1.0, 0.0] + + with pytest.raises(ValueError, match="has no embedding generator.*generate_vectors=False"): + await MockCollection().upsert([Record("missing-generator", "text", [1.0, 0.0])]) + + +async def test_collection_crud_preserves_single_and_batch_shapes() -> None: + collection = MockCollection(embedding_generator=MockEmbeddingClient()) + await collection.ensure_collection_exists() + + first_keys = await collection.upsert([Record("one", "first", "first")]) + keys = await collection.upsert([ + Record("two", "second", "second"), + Record("three", "third", "third"), + ]) + one = await collection.get(["one"]) + many = await collection.get(["one", "two"], include_vectors=True) + filtered = await collection.get(top=1) + + assert first_keys == ["one"] + assert keys == ["two", "three"] + assert one == [Record("one", "first")] + assert many == [ + Record("one", "first", [5.0, 0.5]), + Record("two", "second", [6.0, 0.5]), + ] + assert filtered == [Record("one", "first")] + + await collection.delete(["one", "two"]) + assert await collection.get(["one", "two"]) == [] + + +async def test_collection_wraps_connector_errors() -> None: + collection = MockCollection() + collection.fail_upsert = True + + with pytest.raises(IntegrationException, match="store unavailable"): + await collection.upsert([Record("one", "hello")], generate_vectors=False) + + +async def test_collection_get_without_keys_lists_records() -> None: + assert await MockCollection().get() == [] + + +async def test_collection_crud_rejects_singular_ordinary_inputs() -> None: + collection = MockCollection() + + with pytest.raises(TypeError, match="records must be a sequence"): + await cast(Any, collection.upsert)(Record("one", "hello")) + with pytest.raises(TypeError, match="keys must be a sequence"): + await collection.get("one") + with pytest.raises(TypeError, match="keys must be a sequence"): + await collection.delete("one") + + +async def test_vector_search_generates_query_vector_and_filters_threshold() -> None: + embedding_client = MockEmbeddingClient() + collection = MockCollection(embedding_generator=embedding_client) + await collection.upsert([ + Record("one", "first", "first"), + Record("two", "second", "second"), + ]) + + results = await collection.search( + "find this", + score_threshold=0.5, + ) + responses = [response async for response in results] + + assert results.metadata == {"mock_count": 2} + assert embedding_client.values == ["find this"] + assert collection.last_search_vector == [9.0, 0.5] + assert responses[0]["record"].id == "one" + assert responses[0]["score"] == 0.9 + assert len(responses) == 1 + + +async def test_keyword_hybrid_search_uses_single_search_method() -> None: + collection = MockCollection() + + await collection.search("words", search_type="keyword_hybrid") + + assert collection.last_search_type == "keyword_hybrid" + + +async def test_vector_search_validates_inputs_and_supported_type() -> None: + collection = MockCollection() + + with pytest.raises(ValueError, match="requires values"): + await cast(Any, collection.search)() + + class VectorOnlyCollection(MockCollection): + supported_search_types: ClassVar[set[SearchType]] = {"vector"} + + with pytest.raises(NotImplementedError, match="not supported"): + await VectorOnlyCollection().search("words", search_type="keyword_hybrid") + + +async def test_vector_search_requires_explicit_distance_for_score_threshold() -> None: + collection = MockCollection() + collection.definition = VectorStoreCollectionDefinition( + [ + VectorStoreField("key", name="id", type_="str"), + VectorStoreField("vector", name="vector", type_="float", dimensions=2), + ], + collection_name="records", + ) + + with pytest.raises(ValueError, match="explicit distance"): + await collection.search(vector=[1.0, 0.0], score_threshold=0.5) + + +async def test_vector_search_wraps_embedding_failures() -> None: + class FailingEmbeddingClient(MockEmbeddingClient): + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[list[float]]: + raise RuntimeError("embedding unavailable") + + collection = MockCollection(embedding_generator=FailingEmbeddingClient()) + + with pytest.raises(IntegrationException, match="embedding unavailable"): + await collection.search("query") + + +def test_vector_search_builds_connector_filter() -> None: + collection = MockCollection() + + assert collection._build_filter("lambda record: record.category == 'travel'") == "record.category == 'travel'" + assert collection._build_filter([ + "lambda record: record.category == 'travel'", + "lambda record: record.id != 'ignored'", + ]) == ["record.category == 'travel'", "record.id != 'ignored'"] + + +async def test_vector_search_passes_translated_filter_to_connector() -> None: + collection = MockCollection() + + await collection.search( + "query", + filter="lambda record: record.category == 'travel'", + ) + + assert collection.last_search_filter == "record.category == 'travel'" + + +def test_vector_search_rejects_filter_without_lambda() -> None: + with pytest.raises(ValueError, match="No lambda"): + MockCollection()._build_filter("record.category == 'travel'") + + +async def test_create_search_tool_returns_mapped_results() -> None: + collection = MockCollection() + collection.records["one"] = {"record_id": "one", "body": "first", "vector": [1.0, 0.0]} + tool = create_vector_search_tool( + collection, + name="search_records", + approval_mode="always_require", + top=1, + result_mapper=lambda response: f"{response['record'].id}:{response['score']}", + ) + + result = await tool(query="first") + + assert tool.name == "search_records" + assert tool.approval_mode == "always_require" + assert len(result) == 1 + assert result[0].text == "one:0.9" + + +async def test_create_search_tool_supports_declared_filter_parameters() -> None: + collection = MockCollection() + collection.records["one"] = {"record_id": "one", "body": "first", "vector": [1.0, 0.0]} + tool = create_vector_search_tool( + collection, + parameters={ + "type": "object", + "properties": { + "query": {"type": "string", "description": "The search query."}, + "category": {"type": "string", "description": "The category to match."}, + "top": { + "type": "integer", + "description": "The maximum number of results.", + "maximum": 5, + }, + "skip": { + "type": "integer", + "description": "The number of results to skip.", + "maximum": 10, + }, + }, + "required": ["query", "category"], + "additionalProperties": False, + }, + ) + + await tool(query="first", category="travel", top=1, skip=2) + + assert set(tool.parameters()["properties"]) == {"query", "category", "top", "skip"} + assert collection.last_search_filter == "record.category == 'travel'" + assert collection.last_search_top == 1 + assert collection.last_search_skip == 2 + + +def test_create_search_tool_validates_custom_schema() -> None: + collection = MockCollection() + + with pytest.raises(ValueError, match="required string"): + create_vector_search_tool(collection, parameters={"type": "object", "properties": {}}) + with pytest.raises(ValueError, match="required string"): + create_vector_search_tool( + collection, + parameters={ + "type": "object", + "properties": {"query": {"type": "integer"}}, + "required": ["query"], + }, + ) + with pytest.raises(ValueError, match="declare an integer maximum"): + create_vector_search_tool( + collection, + parameters={ + "type": "object", + "properties": { + "query": {"type": "string"}, + "top": {"type": "integer"}, + }, + "required": ["query"], + }, + ) + + +async def test_create_search_tool_enforces_paging_limits_and_result_cap() -> None: + collection = MockCollection() + collection.raw_search_results = [ + {"record": {"record_id": str(index), "body": f"record {index}"}, "score": 0.9} for index in range(3) + ] + tool = create_vector_search_tool( + collection, + top=2, + parameters={ + "type": "object", + "properties": { + "query": {"type": "string"}, + "top": {"type": "integer", "maximum": 2}, + "skip": {"type": "integer", "maximum": 4}, + }, + "required": ["query"], + }, + ) + + results = await tool(query="records", top=2, skip=4) + assert len(results) == 2 + + with pytest.raises(ValueError, match="top must not exceed"): + await tool(query="records", top=3) + with pytest.raises(ValueError, match="skip must not exceed"): + await tool(query="records", skip=5) + + +async def test_create_search_tool_supports_multimodal_results() -> None: + collection = MockCollection() + collection.records["one"] = {"record_id": "one", "body": "first", "vector": [1.0, 0.0]} + tool = create_vector_search_tool( + collection, + top=1, + result_mapper=lambda response: [ + Content.from_text(response["record"].text), + Content.from_uri("https://example.com/result.png", media_type="image/png"), + ], + ) + + result = await tool.invoke(arguments={"query": "first"}) + + assert [content.type for content in result] == ["text", "uri"] + + +async def test_create_search_tool_uses_msgspec_for_default_result_mapping() -> None: + collection = MockCollection() + collection.records["one"] = {"record_id": "one", "body": "first", "vector": [1.0, 0.0]} + + result = await create_vector_search_tool(collection, top=1)(query="first") + + assert result[0].text is not None + decoded = msgspec.json.decode(result[0].text) + assert set(create_vector_search_tool(collection).parameters()["properties"]) == {"query"} + assert decoded["record"]["id"] == "one" + assert decoded["score"] == 0.9 + + +async def test_create_search_tool_defers_unsupported_type_to_search() -> None: + class VectorOnlyCollection(MockCollection): + supported_search_types: ClassVar[set[SearchType]] = {"vector"} + + tool = create_vector_search_tool(VectorOnlyCollection(), search_type="keyword_hybrid") + with pytest.raises(NotImplementedError, match="not supported"): + await tool(query="query") + + +def test_search_protocol_and_tool_factory_only_require_search() -> None: + class SearchOnly: + async def search( + self, + values: Any, + *, + search_type: SearchType = "vector", + vector: Sequence[float | int] | None = None, + filter: Any = None, + top: int = 3, + skip: int = 0, + include_vectors: bool = False, + vector_property_name: str | None = None, + additional_property_name: str | None = None, + score_threshold: float | None = None, + operation_options: Mapping[str, Any] | None = None, + ) -> SearchResults[SearchResponse[Record]]: + return SearchResults([]) + + search = SearchOnly() + assert isinstance(cast(Any, search), SupportsVectorSearch) + assert create_vector_search_tool(cast(SupportsVectorSearch[Record], search)).name == "search" + + +def test_collection_satisfies_vector_protocols() -> None: + collection = MockCollection() + + assert isinstance(collection, SupportsVectorUpsert) + assert isinstance(collection, SupportsVectorSearch) + + +async def test_vector_store_collection_lifecycle_helpers() -> None: + collection = MockCollection() + store = MockStore(collection) + + assert not await store.collection_exists("records") + await collection.ensure_collection_exists() + assert await store.collection_exists("records") + await store.ensure_collection_deleted("records") + assert not await store.collection_exists("records") + + +def test_search_response_holds_record_and_score() -> None: + record = Record("one", "hello") + response = SearchResponse(record=record, score=0.75) + + assert response["record"] is record + assert response["score"] == 0.75 + + +def test_deserialization_rejects_non_mapping_store_records() -> None: + handler = VectorStoreRecordHandler(Record) + + with pytest.raises(TypeError, match="must be mappings"): + handler.deserialize(object()) + + +def test_additional_field_and_definition_validation_paths() -> None: + with pytest.raises(ValueError, match="Unknown vector store field type"): + cast(Any, VectorStoreField)("unknown") + with pytest.raises(ValueError, match="must not be empty"): + VectorStoreCollectionDefinition([VectorStoreField("key")]) + with pytest.raises(ValueError, match="storage names must be unique"): + VectorStoreCollectionDefinition([ + VectorStoreField("key", name="id", storage_name="same"), + VectorStoreField("data", name="text", storage_name="same"), + ]) + + definition = cast(VectorStoreCollectionDefinition, vars(Record)["__vectorstoremodel_definition__"]) + assert definition.try_get_vector_field("vector") is definition.vector_fields[0] + assert definition.try_get_vector_field("missing") is None + + +async def test_default_codecs_cover_pydantic_plain_and_unsupported_models() -> None: + @vectorstoremodel + class PydanticRecord(BaseModel): + id: Annotated[str, VectorStoreField("key")] + + @vectorstoremodel + class PlainRecord: + id: Annotated[str, VectorStoreField("key")] + + def __init__(self, id: str) -> None: + self.id = id + + @vectorstoremodel + class SlottedRecord: + __slots__ = ("id",) + id: Annotated[str, VectorStoreField("key")] + + def __init__(self, id: str) -> None: + self.id = id + + assert await VectorStoreRecordHandler(PydanticRecord).serialize(PydanticRecord(id="one")) == {"id": "one"} + assert await VectorStoreRecordHandler(PlainRecord).serialize(PlainRecord("one")) == {"id": "one"} + with pytest.raises(NotImplementedError, match="SlottedRecord"): + await VectorStoreRecordHandler(SlottedRecord).serialize(SlottedRecord("one")) + + +def test_vectorstoremodel_rejects_unresolvable_or_missing_annotations() -> None: + class UnresolvableRecord: + __annotations__ = {"id": "MissingRecordType"} + + class EmptyRecord: + pass + + with pytest.raises(ValueError, match="Unable to resolve"): + vectorstoremodel(UnresolvableRecord) + with pytest.raises(ValueError, match="at least one annotated field"): + vectorstoremodel(EmptyRecord) + + +def test_registration_is_idempotent_and_rejects_changed_codecs() -> None: + @dataclass + class RegisteredRecord: + id: str + + definition = VectorStoreCollectionDefinition([VectorStoreField("key", name="id")]) + + def encoder(record: RegisteredRecord) -> Mapping[str, Any]: + return {"id": record.id} + + def decoder(record: Mapping[str, Any]) -> RegisteredRecord: + return RegisteredRecord(cast(str, record["id"])) + + register_vectorstoremodel(RegisteredRecord, definition=definition, encoder=encoder, decoder=decoder) + register_vectorstoremodel(RegisteredRecord, definition=definition, encoder=encoder, decoder=decoder) + + with pytest.raises(ValueError, match="another encoder"): + register_vectorstoremodel( + RegisteredRecord, + definition=definition, + encoder=lambda record: {"id": record.id}, + decoder=decoder, + ) + with pytest.raises(ValueError, match="another decoder"): + register_vectorstoremodel( + RegisteredRecord, + definition=definition, + encoder=encoder, + decoder=lambda record: RegisteredRecord(cast(str, record["id"])), + ) + + +def test_record_handler_requires_registered_models_or_explicit_dict_definitions() -> None: + class UnregisteredRecord: + pass + + with pytest.raises(ValueError, match="explicit"): + VectorStoreRecordHandler(dict) + with pytest.raises(ValueError, match="must be registered"): + VectorStoreRecordHandler(UnregisteredRecord) + + other_definition = VectorStoreCollectionDefinition([VectorStoreField("key", name="other_id")]) + with pytest.raises(ValueError, match="another definition"): + VectorStoreRecordHandler(Record, definition=other_definition) + + +async def test_serialization_shape_and_embedding_failures() -> None: + definition = VectorStoreCollectionDefinition([ + VectorStoreField("key", name="id", storage_name="record_id"), + VectorStoreField("data", name="text", storage_name="body"), + ]) + dict_handler = VectorStoreRecordHandler(dict, definition=definition) + assert await dict_handler.serialize({"record_id": "one", "body": "hello"}) == { + "record_id": "one", + "body": "hello", + } + with pytest.raises(TypeError, match="must serialize to mappings"): + await dict_handler.serialize(cast(Any, 1)) + + collection = MockCollection(embedding_generator=MockEmbeddingClient()) + with pytest.raises(ValueError, match="value is missing"): + await collection.serialize(Record("one", "hello")) + + class EmptyEmbeddingClient(MockEmbeddingClient): + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[list[float]]: + return GeneratedEmbeddings() + + with pytest.raises(IntegrationInvalidResponseException, match="returned 0 vectors"): + await MockCollection(embedding_generator=EmptyEmbeddingClient()).serialize(Record("one", "hello", "embed")) + + assert dict_handler.deserialize(None) is None + + +async def test_array_like_generated_embeddings_are_normalized() -> None: + class ArrayLike: + def tolist(self) -> list[float]: + return [0.1, 0.2] + + class ArrayEmbeddingClient(MockEmbeddingClient): + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[Any]: + return GeneratedEmbeddings([Embedding(vector=ArrayLike()) for _ in values]) + + collection = MockCollection(embedding_generator=ArrayEmbeddingClient()) + serialized = await collection.serialize(Record("one", "hello", "embed")) + assert serialized["vector"] == [0.1, 0.2] + + results = await collection.search("query") + assert collection.last_search_vector == [0.1, 0.2] + assert [result async for result in results] == [] + + +async def test_collection_operation_error_boundaries_and_context_manager() -> None: + collection = MockCollection() + async with collection as entered: + assert entered is collection + + collection.upsert_error = IntegrationException("known upsert failure") + with pytest.raises(IntegrationException, match="known upsert failure"): + await collection.upsert([Record("one", "hello")], generate_vectors=False) + collection.upsert_error = None + collection.upsert_keys = [] + with pytest.raises(IntegrationInvalidResponseException, match="Expected 1 upserted keys"): + await collection.upsert([Record("one", "hello")], generate_vectors=False) + + collection.get_error = RuntimeError("get failure") + with pytest.raises(IntegrationException, match="get failure"): + await collection.get(["one"]) + collection.get_error = None + collection.delete_error = IntegrationException("known delete failure") + with pytest.raises(IntegrationException, match="known delete failure"): + await collection.delete(["one"]) + collection.delete_error = RuntimeError("delete failure") + with pytest.raises(IntegrationException, match="delete failure"): + await collection.delete(["one"]) + + class FailingEmbeddingClient(MockEmbeddingClient): + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[list[float]]: + raise RuntimeError("embedding down") + + with pytest.raises(IntegrationException, match="embedding down"): + await MockCollection(embedding_generator=FailingEmbeddingClient()).upsert([Record("one", "hello", "embed")]) + + +async def test_vector_store_context_and_missing_collection_delete() -> None: + collection = MockCollection() + store = MockStore(collection) + + async with store as entered: + assert entered is store + await store.ensure_collection_deleted("missing") + assert not collection.created + + +async def test_additional_search_validation_and_error_boundaries() -> None: + collection = MockCollection() + with pytest.raises(ValueError, match="Unknown search type"): + await collection.search("query", search_type=cast(Any, "unknown")) + with pytest.raises(ValueError, match="Keyword-hybrid"): + await cast(Any, collection.search)(search_type="keyword_hybrid", vector=[1.0, 0.0]) + with pytest.raises(ValueError, match="was not found"): + await collection.search("query", vector_property_name="missing") + + collection.search_error = IntegrationException("known search failure") + with pytest.raises(IntegrationException, match="known search failure"): + await collection.search("query") + collection.search_error = RuntimeError("search failure") + with pytest.raises(IntegrationException, match="search failure"): + await collection.search("query") + + +async def test_search_embedding_and_result_conversion_failures() -> None: + class EmptyEmbeddingClient(MockEmbeddingClient): + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[list[float]]: + return GeneratedEmbeddings() + + class StringEmbeddingClient(MockEmbeddingClient): + async def get_embeddings( + self, + values: Sequence[Any], + *, + options: EmbeddingGenerationOptions | None = None, + ) -> GeneratedEmbeddings[Any]: + return GeneratedEmbeddings([Embedding(vector="invalid")]) + + with pytest.raises(IntegrationInvalidResponseException, match="returned 0 vectors"): + await MockCollection(embedding_generator=EmptyEmbeddingClient()).search("query") + with pytest.raises(TypeError, match="unsupported vector type"): + await MockCollection(embedding_generator=StringEmbeddingClient()).search("query") + + collection = MockCollection() + collection.raw_search_results = [{"record": None, "score": 0.9}] + results = await collection.search(vector=[1.0, 0.0]) + assert [result async for result in results] == [] + + collection.raw_search_results = [{"record": [{"record_id": "one", "body": "hello"}], "score": 0.9}] + results = await collection.search(vector=[1.0, 0.0]) + with pytest.raises(IntegrationInvalidResponseException, match="exactly one record"): + _ = [result async for result in results] + + collection.raw_search_results = [object()] + results = await collection.search(vector=[1.0, 0.0]) + with pytest.raises(IntegrationInvalidResponseException, match="result conversion failed"): + _ = [result async for result in results] + + async def failing_results() -> AsyncIterable[Any]: + yield {"record": {"record_id": "one", "body": "hello"}, "score": 0.9} + raise RuntimeError("stream disconnected") + + collection.raw_search_results = failing_results() + results = await collection.search(vector=[1.0, 0.0]) + with pytest.raises(IntegrationException, match="iteration failed.*stream disconnected"): + _ = [result async for result in results] + + +async def test_scoreless_results_remain_when_threshold_cannot_be_applied() -> None: + collection = MockCollection() + collection.raw_search_results = [{"record": {"record_id": "one", "body": "hello"}, "score": None}] + + results = await collection.search(vector=[1.0, 0.0], score_threshold=0.5) + + responses = [result async for result in results] + assert len(responses) == 1 + assert responses[0]["score"] is None + + +def test_filter_parser_and_default_mapper_edge_paths() -> None: + collection = MockCollection() + assert collection._build_filter(lambda record: record.id == "one") == "record.id == 'one'" + with pytest.raises(ValueError, match="Unable to parse"): + collection._build_filter("lambda record:") + + +async def test_search_tool_filter_mapper_edge_paths() -> None: + collection = MockCollection() + parameters = { + "type": "object", + "properties": { + "query": {"type": "string"}, + "category": {"type": "string"}, + }, + "required": ["query", "category"], + } + tool = create_vector_search_tool( + collection, + parameters=parameters, + filter=["lambda record: record.id != 'ignored'"], + ) + await tool(query="query", category="travel") + assert collection.last_search_filter == ["record.id != 'ignored'", "record.category == 'travel'"] + + invalid_tool = create_vector_search_tool( + collection, + parameters={ + "type": "object", + "properties": {"query": {"type": "string"}, "bad-name": {"type": "string"}}, + "required": ["query"], + }, + ) + with pytest.raises(ValueError, match="cannot be mapped"): + await invalid_tool(query="query", **{"bad-name": "value"}) + with pytest.raises(TypeError, match="'query'.*string"): + await cast(Any, create_vector_search_tool(collection))(query=1) + + +async def test_runtime_operations_mark_vector_store_feature_usage() -> None: + collection = MockCollection() + store = MockStore(collection) + + with patch("agent_framework._vectors.mark_feature_used") as mark_feature_used_mock: + await collection.serialize(Record("one", "hello"), generate_vectors=False) + mark_feature_used_mock.assert_called_with(FeatureIndex.CORE_VECTOR_STORES) + + mark_feature_used_mock.reset_mock() + collection.deserialize({"record_id": "one", "body": "hello", "vector": None}) + mark_feature_used_mock.assert_called_once_with(FeatureIndex.CORE_VECTOR_STORES) + + mark_feature_used_mock.reset_mock() + await collection.upsert([Record("one", "hello")], generate_vectors=False) + mark_feature_used_mock.assert_any_call(FeatureIndex.CORE_VECTOR_STORES) + + mark_feature_used_mock.reset_mock() + await collection.get(["one"]) + mark_feature_used_mock.assert_any_call(FeatureIndex.CORE_VECTOR_STORES) + + mark_feature_used_mock.reset_mock() + await collection.delete(["one"]) + mark_feature_used_mock.assert_called_once_with(FeatureIndex.CORE_VECTOR_STORES) + + mark_feature_used_mock.reset_mock() + await store.collection_exists("records") + mark_feature_used_mock.assert_called_once_with(FeatureIndex.CORE_VECTOR_STORES) + + mark_feature_used_mock.reset_mock() + await collection.search(vector=[1.0, 0.0]) + mark_feature_used_mock.assert_called_once_with(FeatureIndex.CORE_VECTOR_STORES) diff --git a/python/samples/02-agents/vector_stores/README.md b/python/samples/02-agents/vector_stores/README.md new file mode 100644 index 00000000000..92227f0b1f0 --- /dev/null +++ b/python/samples/02-agents/vector_stores/README.md @@ -0,0 +1,56 @@ +# Vector stores + +Vector stores accept multiple model styles so applications can keep the data +representation that already fits their validation, memory, and interoperability +needs. When you own a model, annotate a dataclass, Pydantic model, msgspec +struct, or plain class. When another team or package owns it, register an +explicit definition and codecs. Dictionaries use a collection-specific +definition; DataFrames and other containers can convert to row dictionaries +before calling the batch API. + +No database or credentials are needed for these examples. + +| File | Demonstrates | +|------|--------------| +| [`vector_store_models.py`](vector_store_models.py) | Choosing among owned models, third-party model registration, and loose dictionary definitions. | +| [`optimized_data_formats.py`](optimized_data_formats.py) | Keeping NumPy vector fields and adapting pandas DataFrames to the batch record API. | + +The first section shows the two equivalent custom-codec registration forms. +`@vectorstoremodel` derives the definition from annotations and registers it; +`register_vectorstoremodel` accepts an externally constructed definition. +Both produce the same internal model registration. + +The sample order is informed by a small benchmark on Apple Silicon with +CPython 3.13. Each benchmark model had the same `id`, `text`, and `vector` +fields. Results are medians of seven warmed runs: + +| Model style | 3-element vector | 1,566-element vector | +|-------------|-----------------:|---------------------:| +| Custom codecs | 2.60 μs | 7.01 μs | +| Dictionary | 2.39 μs | 6.60 μs | +| Plain class | 3.46 μs | 12.26 μs | +| msgspec `Struct` | 3.03 μs | 16.61 μs | +| Dataclass | 3.17 μs | 16.81 μs | +| Pydantic | 5.52 μs | 37.32 μs | + +These results measure only the framework's internal record conversion path. +They do not include database SDK conversion, network I/O, +embedding generation, validation complexity, nested fields, alternate vector +representations, or memory allocation. Custom codecs are especially favorable +here because the benchmark codec returns the existing vector reference rather +than copying it. The middle ordering also changes with vector size, so treat +these timings as illustrative data, not a recommendation or performance +guarantee. + +Array-like vector values, including NumPy arrays, are serialized through their +`tolist()` method without making NumPy a core dependency. If a model must +restore a NumPy array instead of a Python list, pass a custom `decoder` to +`@vectorstoremodel` or `register_vectorstoremodel` and call `numpy.array` or +`numpy.asarray` there. + +Run the sample from the `python` directory: + +```bash +uv run samples/02-agents/vector_stores/vector_store_models.py +uv run samples/02-agents/vector_stores/optimized_data_formats.py +``` diff --git a/python/samples/02-agents/vector_stores/optimized_data_formats.py b/python/samples/02-agents/vector_stores/optimized_data_formats.py new file mode 100644 index 00000000000..a80693d5e5b --- /dev/null +++ b/python/samples/02-agents/vector_stores/optimized_data_formats.py @@ -0,0 +1,121 @@ +# /// script +# requires-python = ">=3.10" +# dependencies = [ +# "agent-framework-core", +# "numpy>=2,<3", +# "pandas>=2,<4", +# ] +# +# [tool.uv.sources] +# agent-framework-core = { path = "../../../packages/core" } +# /// + +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +# Run with: uv run samples/02-agents/vector_stores/optimized_data_formats.py +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Annotated, Any, cast + +import numpy as np +import pandas as pd # pyright: ignore[reportMissingImports] +from agent_framework import ( + VectorStoreCollectionDefinition, + VectorStoreField, + vectorstoremodel, +) +from numpy.typing import NDArray + +"""This sample demonstrates optimized vector data formats. + +When optimized formats already fit the application, Agent Framework should not +force conversion at the model boundary. NumPy arrays reduce the in-memory +footprint of large vectors, while pandas DataFrames can preserve an existing +tabular data pipeline. + +NumPy vectors are encoded through ``tolist()`` without making NumPy a core +dependency. A model decoder restores the array after retrieval. DataFrames are +converted to ordinary row dictionaries before using the vector store batch API +and reconstructed after retrieval. Agent Framework does not need +container-specific behavior or a pandas dependency. + +These formats are choices, not requirements. A plain class or other application +model may be simpler and entirely appropriate. For all standard model and +third-party registration options, see +[vector_store_models.py](vector_store_models.py). +""" + +DIMENSIONS = 1566 + + +# 1. Use a NumPy array as a vector field. +def decode_numpy_record(record: Mapping[str, Any]) -> NumpyRecord: + """Restore a NumPy vector after storage returned an ordinary list.""" + return NumpyRecord( + record_id=cast(str, record["record_id"]), + vector=np.asarray(record["vector"], dtype=np.float32), + ) + + +@vectorstoremodel(collection_name="numpy-records", decoder=decode_numpy_record) +@dataclass +class NumpyRecord: + record_id: Annotated[str, VectorStoreField("key")] + vector: Annotated[ + NDArray[np.float32], + VectorStoreField("vector", dimensions=DIMENSIONS, type_="float"), + ] + + +# 2. Convert a pandas DataFrame to and from ordinary row dictionaries. +dataframe_definition = VectorStoreCollectionDefinition( + [ + VectorStoreField("key", name="id"), + VectorStoreField("data", name="text", is_full_text_indexed=True), + VectorStoreField("vector", name="vector", dimensions=3), + ], + collection_name="dataframe-records", +) + + +def main() -> None: + """Convert NumPy and DataFrame values at the vector store boundary.""" + numpy_vector = np.arange(DIMENSIONS, dtype=np.float32) / np.float32(DIMENSIONS) + serialized_numpy = {"record_id": "numpy-1", "vector": numpy_vector.tolist()} + restored_numpy = decode_numpy_record(serialized_numpy) + + print(f"Serialized vector type: {type(serialized_numpy['vector']).__name__}") + print(f"Restored vector type: {type(restored_numpy.vector).__name__} ({restored_numpy.vector.dtype})") + + frame = pd.DataFrame({ + "id": ["one", "two"], + "text": ["First record", "Second record"], + "vector": [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]], + }) + dataframe_rows = cast(list[dict[str, Any]], frame.to_dict(orient="records")) + restored_frame = pd.DataFrame.from_records(dataframe_rows) + + print(f"Rows passed to batch upsert: {dataframe_rows}") + print("Restored DataFrame:") + print(restored_frame) + + +if __name__ == "__main__": + main() + + +""" +Sample output: +Serialized vector type: list +Restored vector type: ndarray (float32) +Rows passed to batch upsert: [ + {'id': 'one', 'text': 'First record', 'vector': [0.1, 0.2, 0.3]}, + {'id': 'two', 'text': 'Second record', 'vector': [0.4, 0.5, 0.6]} +] +Restored DataFrame: + id text vector +0 one First record [0.1, 0.2, 0.3] +1 two Second record [0.4, 0.5, 0.6] +""" diff --git a/python/samples/02-agents/vector_stores/vector_store_models.py b/python/samples/02-agents/vector_stores/vector_store_models.py new file mode 100644 index 00000000000..d0d79575198 --- /dev/null +++ b/python/samples/02-agents/vector_stores/vector_store_models.py @@ -0,0 +1,174 @@ +# Copyright (c) Microsoft. All rights reserved. + +from __future__ import annotations + +# Run with: uv run samples/02-agents/vector_stores/vector_store_models.py +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Annotated, Any, cast + +import msgspec +from agent_framework import ( + VectorStoreCollectionDefinition, + VectorStoreField, + register_vectorstoremodel, + vectorstoremodel, +) +from pydantic import BaseModel + +"""This sample demonstrates the choices for defining vector store models. + +When you own the model, use the representation that best fits the rest of your +application: a dataclass, Pydantic model, msgspec struct, plain class, or +dictionary. ``@vectorstoremodel`` adds vector store metadata without requiring +the application to adopt one specific modeling library. + +When another package or team owns the model, adapt it instead of rewriting it. +Use ``register_vectorstoremodel`` with an explicit definition and codecs for a +model type, or pass a loose collection definition directly for dictionaries +and other schema-less records. + +For NumPy vectors and DataFrame row containers, see +[optimized_data_formats.py](optimized_data_formats.py). + +The examples are ordered using indicative serialization measurements, but the +right choice depends on validation, memory, interoperability, and the broader +application, not only handler round-trip speed. +""" + + +# 1. Custom codecs can be registered in two equivalent ways. +# The decorator creates the field definition from annotations and registers it with the codecs. +def encode_legacy_faq(faq: LegacyFaq) -> Mapping[str, Any]: + """Convert a legacy FAQ to logical vector model fields.""" + return {"id": faq.faq_number, "question": faq.prompt} + + +def decode_legacy_faq(record: Mapping[str, Any]) -> LegacyFaq: + """Restore a legacy FAQ from logical vector model fields.""" + return LegacyFaq(faq_number=cast(str, record["faq_number"]), prompt=cast(str, record["prompt"])) + + +@vectorstoremodel( + collection_name="legacy-faqs", + encoder=encode_legacy_faq, + decoder=decode_legacy_faq, +) +@dataclass +class LegacyFaq: + faq_number: Annotated[str, VectorStoreField("key", storage_name="id")] + prompt: Annotated[str, VectorStoreField("data", storage_name="question")] + + +# The helper performs the same registration when the definition is supplied separately. +@dataclass +class LegacyArticle: + article_id: int + heading: str + + +def encode_legacy_article(article: LegacyArticle) -> Mapping[str, Any]: + """Convert a legacy article to logical vector model fields.""" + return {"id": str(article.article_id), "title": article.heading} + + +def decode_legacy_article(record: Mapping[str, Any]) -> LegacyArticle: + """Restore a legacy article from logical vector model fields.""" + return LegacyArticle(article_id=int(record["id"]), heading=cast(str, record["title"])) + + +legacy_definition = VectorStoreCollectionDefinition( + [ + VectorStoreField("key", name="id", storage_name="article_id"), + VectorStoreField("data", name="title", storage_name="heading"), + ], + collection_name="legacy-articles", +) +register_vectorstoremodel( + LegacyArticle, + definition=legacy_definition, + encoder=encode_legacy_article, + decoder=decode_legacy_article, +) + + +# 2. Plain dictionaries can be used as models; in that case, we just need the collection-specific definition. +dictionary_definition = VectorStoreCollectionDefinition( + [ + VectorStoreField("key", name="id"), + VectorStoreField("data", name="text"), + VectorStoreField("vector", name="vector", dimensions=3), + ], + collection_name="dictionary-records", +) + + +# 3. Plain classes use their annotated constructor parameters. +@vectorstoremodel(collection_name="notes") +class Note: + def __init__( + self, + note_id: Annotated[str, VectorStoreField("key")], + text: Annotated[str, VectorStoreField("data")], + ) -> None: + self.note_id = note_id + self.text = text + + +# 4. msgspec structs use the default registered codec. +@vectorstoremodel(collection_name="documents") +class Document(msgspec.Struct): + document_id: Annotated[str, VectorStoreField("key")] + title: Annotated[str, VectorStoreField("data")] + vector: Annotated[list[float] | None, VectorStoreField("vector", dimensions=3)] = None + + +# 5. Dataclasses use the default registered codec. +@vectorstoremodel(collection_name="hotels") +@dataclass +class Hotel: + hotel_id: Annotated[str, VectorStoreField("key")] + name: Annotated[str, VectorStoreField("data", is_indexed=True)] + description: Annotated[ + str | list[float] | None, + VectorStoreField("vector", dimensions=3, distance_function="cosine_similarity"), + ] = None + + +# 6. Pydantic models provide validation with additional round-trip cost. +@vectorstoremodel(collection_name="products") +class Product(BaseModel): + product_id: Annotated[str, VectorStoreField("key")] + name: Annotated[str, VectorStoreField("data", is_full_text_indexed=True)] + vector: Annotated[list[float] | None, VectorStoreField("vector", dimensions=3)] = None + + +def main() -> None: + """Inspect model definitions and registration choices.""" + model_definitions = ( + ("LegacyFaq", cast(VectorStoreCollectionDefinition, vars(LegacyFaq)["__vectorstoremodel_definition__"])), + ("LegacyArticle", legacy_definition), + ("dict", dictionary_definition), + ("Note", cast(VectorStoreCollectionDefinition, vars(Note)["__vectorstoremodel_definition__"])), + ("Document", cast(VectorStoreCollectionDefinition, vars(Document)["__vectorstoremodel_definition__"])), + ("Hotel", cast(VectorStoreCollectionDefinition, vars(Hotel)["__vectorstoremodel_definition__"])), + ("Product", cast(VectorStoreCollectionDefinition, vars(Product)["__vectorstoremodel_definition__"])), + ) + for model_name, definition in model_definitions: + print(f"{model_name}: collection={definition.collection_name}, fields={definition.names}") + + +if __name__ == "__main__": + main() + + +""" +Sample output: +LegacyFaq: collection=legacy-faqs, fields=['faq_number', 'prompt'] +LegacyArticle: collection=legacy-articles, fields=['id', 'title'] +dict: collection=dictionary-records, fields=['id', 'text', 'vector'] +Note: collection=notes, fields=['note_id', 'text'] +Document: collection=documents, fields=['document_id', 'title', 'vector'] +Hotel: collection=hotels, fields=['hotel_id', 'name', 'description'] +Product: collection=products, fields=['product_id', 'name', 'vector'] +""" diff --git a/python/samples/AGENTS.md b/python/samples/AGENTS.md index 7250e424f62..739329a1250 100644 --- a/python/samples/AGENTS.md +++ b/python/samples/AGENTS.md @@ -10,7 +10,8 @@ python/samples/ ├── 01-get-started/ # Progressive tutorial (steps 01–07) ├── 02-agents/ # Deep-dive concept samples │ ├── tools/ # Tool patterns (function, approval, schema, etc.) -│ ├── middleware/ # One file per middleware concept +│ ├── vector_stores/ # Vector model schemas and registration +│ ├── middleware/ # One file per middleware concept │ ├── conversations/ # Thread, storage, suspend/resume │ ├── providers/ # One sub-folder per provider (azure_ai/, openai/, etc.) │ ├── context_providers/ # Memory & context injection