diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 3a9021c2b94..237eb463518 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -27,7 +27,7 @@ Describe key changes, mention related issues or motivation for the changes. ### Duplicate and AI-Generated PR Check -- [ ] I have searched existing [open pull requests](../../pulls) and confirmed that no other PR already addresses this issue +- [ ] I have searched existing [open pull requests](https://github.com/agno-agi/agno/pulls) and confirmed that no other PR already addresses this issue - [ ] If a similar PR exists, I have explained below why this PR is a better approach - [ ] Check if this PR was entirely AI-generated (by Copilot, Claude Code, Cursor, etc.) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 92504ea9169..3c128a64db1 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -148,4 +148,4 @@ Message us on [Discord](https://discord.gg/4MtYHHrgA8) if you have any questions ## 📝 License -This project is licensed under the terms of the [Apache-2.0 license](/LICENSE) +This project is licensed under the terms of the [Apache-2.0 license](LICENSE) diff --git a/cookbook/02_agents/07_knowledge/agentic_rag_with_reranking.py b/cookbook/02_agents/07_knowledge/agentic_rag_with_reranking.py index 22e44ce36c3..51066830b51 100644 --- a/cookbook/02_agents/07_knowledge/agentic_rag_with_reranking.py +++ b/cookbook/02_agents/07_knowledge/agentic_rag_with_reranking.py @@ -41,5 +41,5 @@ # Run Agent # --------------------------------------------------------------------------- if __name__ == "__main__": - knowledge.insert(name="Agno Docs", url="https://docs.agno.com/introduction.md") + knowledge.insert(name="Agno Docs", url="https://docs.agno.com/introduction") agent.print_response("What are Agno's key features?") diff --git a/cookbook/07_knowledge/02_building_blocks/07_knowledge_level_reranking.py b/cookbook/07_knowledge/02_building_blocks/07_knowledge_level_reranking.py new file mode 100644 index 00000000000..626fbee8151 --- /dev/null +++ b/cookbook/07_knowledge/02_building_blocks/07_knowledge_level_reranking.py @@ -0,0 +1,67 @@ +""" +Knowledge-Level Reranking +========================= +A reranker set on Knowledge runs after the vector db returns results, rather than +inside the vector db itself. Two differences follow from that: + +1. It works with any vector db, so the same reranker moves between backends. +2. It widens the fetch, so the reranker chooses from a real pool rather than only + reordering what the vector db already returned. candidate_multiplier (capped by + max_candidates) is set on the reranker itself. + +The widened fetch is what makes ordering strategies possible: a reranker can only +surface a document that was retrieved in the first place. + +See also: 03_reranking.py for vector db level reranking. +""" + +import asyncio + +from agno.agent import Agent +from agno.knowledge.knowledge import Knowledge +from agno.knowledge.reranker.cohere import CohereReranker +from agno.models.openai import OpenAIResponses +from agno.vectordb.qdrant import Qdrant + +# --------------------------------------------------------------------------- +# Setup +# --------------------------------------------------------------------------- + +qdrant_url = "http://localhost:6333" + +knowledge = Knowledge( + vector_db=Qdrant(collection="knowledge_reranking_demo", url=qdrant_url), + reranker=CohereReranker( + # Candidates fetched per requested result, so Cohere can rescue a document that + # plain search ranked outside max_results. Costs that many times the API calls, + # so lower it to 1 to only reorder what the vector db already returned. + candidate_multiplier=3, + # Ceiling on the widened fetch, once the multiplier is above 1. + max_candidates=100, + ), +) + +agent = Agent( + model=OpenAIResponses(id="gpt-5.6-luna"), + knowledge=knowledge, + search_knowledge=True, + markdown=True, +) + + +async def main(): + await knowledge.ainsert( + url="https://agno-public.s3.amazonaws.com/recipes/ThaiRecipes.pdf" + ) + + # Retrieves 25 candidates, reranks them, returns the top 5. + results = await knowledge.asearch("What are some Thai curry dishes?", max_results=5) + print("Reranked results:") + for document in results: + print(f" {document.name}") + + await agent.aprint_response("What are some Thai curry dishes?", stream=True) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/cookbook/07_knowledge/02_building_blocks/08_mmr_diverse_results.py b/cookbook/07_knowledge/02_building_blocks/08_mmr_diverse_results.py new file mode 100644 index 00000000000..be44d1963b8 --- /dev/null +++ b/cookbook/07_knowledge/02_building_blocks/08_mmr_diverse_results.py @@ -0,0 +1,96 @@ +""" +MMR: Diverse, Non-Redundant Results +=================================== +Vector search returns the closest matches to a query, which are often near-duplicates +of each other: five chunks that all say the same thing. MMR (Maximal Marginal +Relevance) picks documents one at a time, discounting each candidate by how similar it +already is to what has been selected. + +lambda_mult controls the tradeoff: +- 1.0 ranks by relevance alone (equivalent to plain vector search) +- 0.5 balances relevance against difference +- 0.0 ranks by difference alone + +MMR needs a pool larger than the number of results requested, which is what the +reranker provides: candidate_multiplier widens the fetch, MMR selects from it, and +max_results are returned. + +MMR reads the embedding on each search result. Not every vector db returns one: +Milvus, MongoDB, Redis and Valkey do not, so MMR raises there rather +than silently returning unreranked results. + +Take the returned order as the result: reranking_score holds the MMR score at the +moment each document was picked, which is not descending, so re-sorting by it discards +the diversity ordering. + +Set a reranker in one place: with one on both Knowledge and the vector db, only the +one on Knowledge is applied and the vector db's is ignored. + +See also: 07_knowledge_level_reranking.py for how the widened fetch works. +""" + +import asyncio + +from agno.agent import Agent +from agno.knowledge.knowledge import Knowledge +from agno.knowledge.reranker.mmr import MMRReranker +from agno.models.openai import OpenAIResponses +from agno.vectordb.qdrant import Qdrant + +# --------------------------------------------------------------------------- +# Setup +# --------------------------------------------------------------------------- + +qdrant_url = "http://localhost:6333" + +knowledge = Knowledge( + vector_db=Qdrant(collection="mmr_demo", url=qdrant_url), + reranker=MMRReranker( + # Relevance against diversity: 1.0 is relevance alone, 0.0 difference alone. + lambda_mult=0.5, + # Candidates fetched per requested result, so MMR has a pool to choose from. + candidate_multiplier=5, + # Ceiling on that widened fetch, whatever max_results is asked for. + max_candidates=100, + ), +) + +agent = Agent( + model=OpenAIResponses(id="gpt-5.6-luna"), + knowledge=knowledge, + markdown=True, +) + + +def show(results, candidates: int) -> None: + """Print a snippet per result: every chunk shares the source file name.""" + print(f"Selected {len(results)} of {candidates} candidates:\n") + for document in results: + snippet = " ".join(document.content.split())[:100] + print(f" - {snippet}...") + print() + + +async def main(): + await knowledge.ainsert( + url="https://agno-public.s3.amazonaws.com/recipes/ThaiRecipes.pdf" + ) + + query = "What are some Thai curry dishes?" + + # Same query without MMR, to compare against. + plain = Knowledge(vector_db=knowledge.vector_db) + candidates = len(await plain.asearch(query, max_results=25)) + + print("\nWithout MMR") + show(await plain.asearch(query, max_results=5), candidates) + + # Retrieves 25 candidates, selects 5 that are relevant but unlike each other. + print("With MMR") + show(await knowledge.asearch(query, max_results=5), candidates) + + await agent.aprint_response("What are some Thai curry dishes?", stream=True) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/cookbook/07_knowledge/02_building_blocks/09_mmr_with_pgvector.py b/cookbook/07_knowledge/02_building_blocks/09_mmr_with_pgvector.py new file mode 100644 index 00000000000..ff6fd9d905c --- /dev/null +++ b/cookbook/07_knowledge/02_building_blocks/09_mmr_with_pgvector.py @@ -0,0 +1,104 @@ +""" +MMR with PgVector +================= +The same diversity selection as 08_mmr_diverse_results.py, against PgVector. + +MMR compares candidates to each other, so it needs the embedding of every search +result. PgVector returns embeddings on search, so MMR works against it directly. + +Setup: + ./cookbook/scripts/run_pgvector.sh + +See also: 08_mmr_diverse_results.py for what lambda_mult controls. +""" + +import asyncio + +from agno.agent import Agent +from agno.knowledge.embedder.openai import OpenAIEmbedder +from agno.knowledge.knowledge import Knowledge +from agno.knowledge.reranker.mmr import MMRReranker +from agno.models.openai import OpenAIResponses +from agno.vectordb.pgvector import PgVector +from agno.vectordb.search import SearchType + +# --------------------------------------------------------------------------- +# Setup +# --------------------------------------------------------------------------- + +db_url = "postgresql+psycopg://ai:ai@localhost:5532/ai" + +knowledge = Knowledge( + vector_db=PgVector( + table_name="mmr_demo", + db_url=db_url, + search_type=SearchType.hybrid, + embedder=OpenAIEmbedder(id="text-embedding-3-small"), + ), + # Runs after PgVector returns candidates. + reranker=MMRReranker( + # Relevance against diversity: 1.0 is relevance alone, 0.0 difference alone. + lambda_mult=0.5, + # Candidates fetched per requested result, so MMR has a pool to choose from. + candidate_multiplier=5, + # Ceiling on that widened fetch, whatever max_results is asked for. + max_candidates=100, + ), +) + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- + +agent = Agent( + model=OpenAIResponses(id="gpt-5.6-luna"), + knowledge=knowledge, + search_knowledge=True, + instructions=[ + "Always search your knowledge base before answering.", + "Include sources in your response.", + ], + markdown=True, +) + +# --------------------------------------------------------------------------- +# Run Demo +# --------------------------------------------------------------------------- + + +def show(results, candidates: int) -> None: + """Print a snippet per result: every chunk shares the source file name.""" + print(f"Selected {len(results)} of {candidates} candidates:\n") + for document in results: + snippet = " ".join(document.content.split())[:100] + print(f" - {snippet}...") + print() + + +if __name__ == "__main__": + + async def main(): + await knowledge.ainsert( + url="https://agno-public.s3.amazonaws.com/recipes/ThaiRecipes.pdf" + ) + + print("\n" + "=" * 60) + print("PgVector hybrid search + MMR") + print("=" * 60 + "\n") + + query = "What are some Thai curry dishes?" + + # Same query without MMR, to compare against. + plain = Knowledge(vector_db=knowledge.vector_db) + candidates = len(await plain.asearch(query, max_results=25)) + + print("Without MMR") + show(await plain.asearch(query, max_results=5), candidates) + + # Retrieves 25 candidates, selects 5 that are relevant but unlike each other. + print("With MMR") + show(await knowledge.asearch(query, max_results=5), candidates) + + await agent.aprint_response("What are some Thai curry dishes?", stream=True) + + asyncio.run(main()) diff --git a/cookbook/07_knowledge/02_building_blocks/10_mmr_with_elasticsearch.py b/cookbook/07_knowledge/02_building_blocks/10_mmr_with_elasticsearch.py new file mode 100644 index 00000000000..4691fc4760a --- /dev/null +++ b/cookbook/07_knowledge/02_building_blocks/10_mmr_with_elasticsearch.py @@ -0,0 +1,110 @@ +""" +MMR with Elasticsearch +====================== +The same diversity selection as 08_mmr_diverse_results.py, against Elasticsearch. + +MMR compares candidates to each other, so it needs the embedding of every search +result. Elasticsearch returns embeddings on search, so MMR works against it directly. + +Setup: + ./cookbook/scripts/run_elasticsearch.sh + +See also: 08_mmr_diverse_results.py for what lambda_mult controls. +""" + +import asyncio + +from agno.agent import Agent +from agno.knowledge.embedder.openai import OpenAIEmbedder +from agno.knowledge.knowledge import Knowledge +from agno.knowledge.reranker.mmr import MMRReranker +from agno.models.openai import OpenAIResponses +from agno.vectordb.elasticsearch import Elasticsearch +from agno.vectordb.search import SearchType + +# --------------------------------------------------------------------------- +# Setup +# --------------------------------------------------------------------------- + +elasticsearch_url = "http://localhost:9200" + +vector_db = Elasticsearch( + index_name="mmr_demo", + url=elasticsearch_url, + search_type=SearchType.hybrid, + embedder=OpenAIEmbedder(id="text-embedding-3-small"), +) + +knowledge = Knowledge( + vector_db=vector_db, + # Runs after Elasticsearch returns candidates. + reranker=MMRReranker( + # Relevance against diversity: 1.0 is relevance alone, 0.0 difference alone. + lambda_mult=0.5, + # Candidates fetched per requested result, so MMR has a pool to choose from. + candidate_multiplier=5, + # Ceiling on that widened fetch, whatever max_results is asked for. + max_candidates=100, + ), +) + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- + +agent = Agent( + model=OpenAIResponses(id="gpt-5.6-luna"), + knowledge=knowledge, + search_knowledge=True, + instructions=[ + "Always search your knowledge base before answering.", + "Include sources in your response.", + ], + markdown=True, +) + +# --------------------------------------------------------------------------- +# Run Demo +# --------------------------------------------------------------------------- + + +def show(results, candidates: int) -> None: + """Print a snippet per result: every chunk shares the source file name.""" + print(f"Selected {len(results)} of {candidates} candidates:\n") + for document in results: + snippet = " ".join(document.content.split())[:100] + print(f" - {snippet}...") + print() + + +if __name__ == "__main__": + + async def main(): + await knowledge.ainsert( + url="https://agno-public.s3.amazonaws.com/recipes/ThaiRecipes.pdf" + ) + + print("\n" + "=" * 60) + print("Elasticsearch hybrid search + MMR") + print("=" * 60 + "\n") + + query = "What are some Thai curry dishes?" + + # Same query without MMR, to compare against. + plain = Knowledge(vector_db=knowledge.vector_db) + candidates = len(await plain.asearch(query, max_results=25)) + + print("Without MMR") + show(await plain.asearch(query, max_results=5), candidates) + + # Retrieves 25 candidates, selects 5 that are relevant but unlike each other. + print("With MMR") + show(await knowledge.asearch(query, max_results=5), candidates) + + await agent.aprint_response("What are some Thai curry dishes?", stream=True) + + # The async client holds an aiohttp session that Python will not close for + # you: skip this and the script exits with an unclosed connector warning. + await vector_db.async_close() + + asyncio.run(main()) diff --git a/cookbook/07_knowledge/09_archive/vector_dbs/llamaindex_db.py b/cookbook/07_knowledge/09_archive/vector_dbs/llamaindex_db.py index ffb8e45f674..25b2c31f6bc 100644 --- a/cookbook/07_knowledge/09_archive/vector_dbs/llamaindex_db.py +++ b/cookbook/07_knowledge/09_archive/vector_dbs/llamaindex_db.py @@ -23,7 +23,7 @@ # Setup # --------------------------------------------------------------------------- data_dir = Path(__file__).parent.parent.parent.joinpath("wip", "data", "paul_graham") -source_url = "https://raw.githubusercontent.com/run-llama/llama_index/main/docs/docs/examples/data/paul_graham/paul_graham_essay.txt" +source_url = "https://raw.githubusercontent.com/run-llama/llama_index/main/docs/examples/data/paul_graham/paul_graham_essay.txt" # --------------------------------------------------------------------------- diff --git a/cookbook/07_knowledge/09_archive/vector_dbs/pgvector_with_bedrock_reranker.py b/cookbook/07_knowledge/09_archive/vector_dbs/pgvector_with_bedrock_reranker.py index 3874deffb9b..29ab3aa35bd 100644 --- a/cookbook/07_knowledge/09_archive/vector_dbs/pgvector_with_bedrock_reranker.py +++ b/cookbook/07_knowledge/09_archive/vector_dbs/pgvector_with_bedrock_reranker.py @@ -84,7 +84,7 @@ # --------------------------------------------------------------------------- def main() -> None: knowledge_cohere.insert( - name="Agno Docs", url="https://docs.agno.com/introduction.md" + name="Agno Docs", url="https://docs.agno.com/introduction" ) _ = knowledge_convenience _ = knowledge_amazon diff --git a/cookbook/07_knowledge/09_archive/vector_dbs/redis_db_with_cohere_reranker.py b/cookbook/07_knowledge/09_archive/vector_dbs/redis_db_with_cohere_reranker.py index d0f1a8f8b18..dee6fa0a5ad 100644 --- a/cookbook/07_knowledge/09_archive/vector_dbs/redis_db_with_cohere_reranker.py +++ b/cookbook/07_knowledge/09_archive/vector_dbs/redis_db_with_cohere_reranker.py @@ -39,7 +39,7 @@ # Run Agent # --------------------------------------------------------------------------- def main() -> None: - knowledge.insert(name="Agno Docs", url="https://docs.agno.com/introduction.md") + knowledge.insert(name="Agno Docs", url="https://docs.agno.com/introduction") agent.print_response("What are Agno's key features?") diff --git a/cookbook/11_memory/integrations/README.md b/cookbook/11_memory/integrations/README.md index ea13fde61df..45a8e4b0331 100644 --- a/cookbook/11_memory/integrations/README.md +++ b/cookbook/11_memory/integrations/README.md @@ -7,6 +7,7 @@ Examples for connecting Agno agents to external memory services. - [`mem0_integration.py`](./mem0_integration.py): Uses Mem0 as an external memory service. - [`memori_integration.py`](./memori_integration.py): Uses Memori for conversation memory. - [`zep_integration.py`](./zep_integration.py): Uses Zep tools for memory context retrieval. +- [`inspeximus_integration.py`](./inspeximus_integration.py): Uses inspeximus for corrections that stay corrected. ## Run diff --git a/cookbook/11_memory/integrations/inspeximus_integration.py b/cookbook/11_memory/integrations/inspeximus_integration.py new file mode 100644 index 00000000000..f1b59ef4059 --- /dev/null +++ b/cookbook/11_memory/integrations/inspeximus_integration.py @@ -0,0 +1,62 @@ +""" +inspeximus Integration +====================== + +Demonstrates memory that stays corrected: a later write to the same key retires the earlier value. +""" + +import os + +from agno.agent import Agent +from agno.models.openai import OpenAIChat +from inspeximus import Inspeximus + +# --------------------------------------------------------------------------- +# Setup +# --------------------------------------------------------------------------- +# Start from an empty store so the example is repeatable. Re-running it against the +# store it left behind would restate a value that revert() had already retired, and +# inspeximus refuses that on purpose: an echo does not un-retire a correction. +if os.path.exists("agno_memory.json"): + os.remove("agno_memory.json") + +memory = Inspeximus("agno_memory.json") + +# A key is what makes the second write RETIRE the first, with no model call and no +# similarity threshold. Without a key, a write is an ordinary appended fact. +memory.remember("The staging database is db-3.internal", key="staging-db") +# ... and the correction, under the same key. +memory.remember("The staging database is db-7.internal", key="staging-db") + +# Only the correction comes back. The retired value stays in the history, and recall +# does not hand it to the agent. +current = [hit["text"] for hit in memory.recall("staging database", k=3)] + + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- +agent = Agent( + model=OpenAIChat(), + instructions=["Answer from the remembered facts you are given."], + dependencies={"memory": "\n".join(current)}, + add_dependencies_to_context=True, + markdown=True, +) + + +# --------------------------------------------------------------------------- +# Run Example +# --------------------------------------------------------------------------- +if __name__ == "__main__": + assert current == ["The staging database is db-7.internal"], current + agent.print_response("Which staging database should I use?") + + # A correction is reversible: revert(key) makes the previous value current again, + # without naming it, and the agent sees the change on the next turn. + memory.revert("staging-db") + restored = [hit["text"] for hit in memory.recall("staging database", k=3)] + assert restored == ["The staging database is db-3.internal"], restored + + agent.dependencies = {"memory": "\n".join(restored)} + agent.print_response("Which staging database should I use?") diff --git a/cookbook/90_models/anthropic/mcp_connector.py b/cookbook/90_models/anthropic/mcp_connector.py index 39a07db7177..a16fef975b2 100644 --- a/cookbook/90_models/anthropic/mcp_connector.py +++ b/cookbook/90_models/anthropic/mcp_connector.py @@ -21,7 +21,7 @@ MCPServerConfiguration( type="url", name="deepwiki", - url="https://mcp.deepwiki.com/sse", + url="https://mcp.deepwiki.com/mcp", ) ], ), diff --git a/cookbook/90_models/yapi/README.md b/cookbook/90_models/yapi/README.md new file mode 100644 index 00000000000..d690d7607ec --- /dev/null +++ b/cookbook/90_models/yapi/README.md @@ -0,0 +1,74 @@ +# Y-API Cookbook + +This cookbook demonstrates how to use Y-API with the Agno framework. Y-API is an +OpenAI-compatible gateway that serves models from several vendors behind a single endpoint, +with org-prefixed model ids (`deepseek/deepseek-v4-flash`, `z-ai/glm-5.3`, `openai/gpt-5.6-sol`, ...). + +> **Prerequisites**: Fork and clone this repository if needed + +## Quick Start + +### 1. Create and activate a virtual environment + +```shell +python3 -m venv ~/.venvs/aienv +source ~/.venvs/aienv/bin/activate +``` + +### 2. Export your `YAPI_API_KEY` + +Get your API key from: https://y-api.bestvirtualgoods.com/app/keys + +```shell +export YAPI_API_KEY=sk-*** +``` + +### 3. Install libraries + +```shell +uv pip install -U openai agno +``` + +### 4. Run basic Agent + +```shell +python cookbook/90_models/yapi/basic.py +``` + +### 5. Run Agent with Tools + +```shell +python cookbook/90_models/yapi/tool_use.py +``` + +## Model Ids + +The endpoint reports its current catalogue at `GET /v1/models`, and the ids are org-prefixed +(`/`). A few examples: + +- `deepseek/deepseek-v4-flash` (default) — the cheapest option, and it supports tool calls +- `deepseek/deepseek-v4-pro` +- `z-ai/glm-5.3`, `z-ai/glm-5.2` +- `moonshotai/kimi-k3` +- `openai/gpt-5.6-sol`, `openai/gpt-5.6-terra`, `openai/gpt-5.6-luna`, `openai/gpt-6-astra` +- `qwen/qwen3.8-flash`, `tencent/hy3`, `xiaomi/mimo-v2.5` + +The catalogue changes, so read it from `/v1/models` rather than hardcoding it. + +> **Note**: The four `openai/gpt-5.6-*` and `openai/gpt-6-astra` ids require an explicit +> `reasoning_effort` when `tools` are sent — without it the endpoint answers +> `Function tools with reasoning_effort are not supported`. If you hit that with an agent, +> pass `reasoning_effort="low"` on the model, or use one of the other ids above. + +## Resources & Support + +### 🔗 Official Links +- [Website](https://y-api.bestvirtualgoods.com) +- [Documentation](https://y-api.bestvirtualgoods.com/docs) +- [Model List](https://y-api.bestvirtualgoods.com/models) +- [Pricing](https://y-api.bestvirtualgoods.com/pricing) +- [Get API Key](https://y-api.bestvirtualgoods.com/app/keys) + +### 📖 API Reference +- **Base URL**: `https://api.y-api.bestvirtualgoods.com/v1` +- **Models Endpoint**: `https://api.y-api.bestvirtualgoods.com/v1/models` diff --git a/cookbook/90_models/yapi/TEST_LOG.md b/cookbook/90_models/yapi/TEST_LOG.md new file mode 100644 index 00000000000..ba5bebec16e --- /dev/null +++ b/cookbook/90_models/yapi/TEST_LOG.md @@ -0,0 +1,3 @@ +# TEST_LOG + +No cookbook tests have been recorded for this directory yet. diff --git a/cookbook/90_models/yapi/basic.py b/cookbook/90_models/yapi/basic.py new file mode 100644 index 00000000000..e6162341340 --- /dev/null +++ b/cookbook/90_models/yapi/basic.py @@ -0,0 +1,40 @@ +""" +Yapi Basic +========== + +Cookbook example for YAPI, OpenAILike model provider. +""" + +import asyncio + +from agno.agent import Agent +from agno.models.yapi import YAPI + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- + +agent = Agent(model=YAPI(id="deepseek/deepseek-v4-flash"), markdown=True) + +# You can also select the provider by string, which resolves to the same class: +# agent = Agent(model="yapi:deepseek/deepseek-v4-flash", markdown=True) + +# --------------------------------------------------------------------------- +# Run Agent +# --------------------------------------------------------------------------- +if __name__ == "__main__": + # --- Sync --- + agent.print_response("Explain quantum computing in simple terms") + + # --- Sync + Streaming --- + agent.print_response("Explain quantum computing in simple terms", stream=True) + + # --- Async --- + asyncio.run(agent.aprint_response("Share a 2 sentence horror story")) + + # --- Async + Streaming --- + asyncio.run( + agent.aprint_response( + "Write a short poem about artificial intelligence", stream=True + ) + ) diff --git a/cookbook/90_models/yapi/tool_use.py b/cookbook/90_models/yapi/tool_use.py new file mode 100644 index 00000000000..d20107d1315 --- /dev/null +++ b/cookbook/90_models/yapi/tool_use.py @@ -0,0 +1,45 @@ +""" +Yapi Tool Use +============= + +Cookbook example for `yapi/tool_use.py`. +""" + +import asyncio + +from agno.agent import Agent +from agno.models.yapi import YAPI +from agno.tools.websearch import WebSearchTools + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- + +agent = Agent( + model=YAPI(id="deepseek/deepseek-v4-flash"), + tools=[WebSearchTools()], + markdown=True, +) + +# --------------------------------------------------------------------------- +# Run Agent +# --------------------------------------------------------------------------- +if __name__ == "__main__": + # --- Sync --- + agent.print_response("What's the latest news about AI?") + + # --- Sync + Streaming --- + agent.print_response("What's the current weather in Tokyo?", stream=True) + + # --- Async --- + asyncio.run( + agent.aprint_response("What is the latest price about BTCUSDT on Binance?") + ) + + # --- Async + Streaming --- + asyncio.run( + agent.aprint_response( + "Search for the latest developments in quantum computing and summarize them", + stream=True, + ) + ) diff --git a/cookbook/91_tools/coding_tools/01_basic_usage.py b/cookbook/91_tools/coding_tools/01_basic_usage.py index a46cde978dc..56b95bfbded 100644 --- a/cookbook/91_tools/coding_tools/01_basic_usage.py +++ b/cookbook/91_tools/coding_tools/01_basic_usage.py @@ -1,18 +1,18 @@ """ CodingTools: Minimal Tools for Coding Agents ============================================= -A single toolkit with 4 core tools (read, edit, write, shell) that lets -an agent perform any coding task. Inspired by the Pi coding agent's -philosophy: a small number of composable tools is more powerful than -many specialized ones. +A single toolkit that lets an agent perform any coding task. Inspired by +the Pi coding agent's philosophy: a small number of composable tools is +more powerful than many specialized ones. Core tools (enabled by default): - read_file: Read files with line numbers and pagination - edit_file: Exact text find-and-replace with diff output - write_file: Create or overwrite files -- run_shell: Execute shell commands with timeout -Exploration tools (opt-in): +Opt-in tools: +- run_shell: Execute shell commands with timeout. Off by default because it + runs arbitrary commands; enable it only for agents you supervise. - grep: Search file contents - find: Search for files by glob pattern - ls: List directory contents @@ -27,7 +27,8 @@ # --------------------------------------------------------------------------- agent = Agent( model=OpenAIResponses(id="gpt-5.2"), - tools=[CodingTools(base_dir=".")], + # run_shell is opt-in; enable it here to let the agent list the directory. + tools=[CodingTools(base_dir=".", enable_run_shell=True)], instructions="You are a coding assistant. Use the coding tools to help the user.", markdown=True, ) diff --git a/cookbook/91_tools/coding_tools/README.md b/cookbook/91_tools/coding_tools/README.md index 84000ee2968..02e21cd3ee1 100644 --- a/cookbook/91_tools/coding_tools/README.md +++ b/cookbook/91_tools/coding_tools/README.md @@ -1,6 +1,6 @@ # CodingTools -A minimal, powerful toolkit for coding agents. Provides 4 core tools and 3 optional exploration tools. +A minimal, powerful toolkit for coding agents. Provides 3 core tools plus opt-in shell and exploration tools. ## Philosophy @@ -15,12 +15,12 @@ Inspired by the Pi coding agent: a small number of composable tools is more powe | `read_file` | Read files with line numbers and pagination | | `edit_file` | Exact text find-and-replace with unified diff output | | `write_file` | Create or overwrite files, auto-creates parent dirs | -| `run_shell` | Execute shell commands with timeout and output truncation | -### Exploration (opt-in) +### Opt-in | Tool | Description | |------|-------------| +| `run_shell` | Execute shell commands with timeout and output truncation. Off by default; runs arbitrary commands, so enable only under supervision and never for untrusted input. `restrict_to_base_dir` limits accidental damage but is not a security sandbox. | | `grep` | Search file contents for a pattern | | `find` | Search for files by glob pattern | | `ls` | List directory contents | @@ -32,7 +32,7 @@ from agno.agent import Agent from agno.models.openai import OpenAIChat from agno.tools.coding import CodingTools -# Core tools only (default) +# Core tools only (read, edit, write) agent = Agent( model=OpenAIChat(id="gpt-5.6-luna"), tools=[CodingTools(base_dir="./workspace")], @@ -55,5 +55,5 @@ agent = Agent( | File | Description | |------|-------------| -| `01_basic_usage.py` | Core 4 tools with a coding agent | +| `01_basic_usage.py` | Core tools plus opt-in run_shell | | `02_all_tools.py` | All 7 tools enabled | diff --git a/cookbook/91_tools/mcp/README.md b/cookbook/91_tools/mcp/README.md index 5419a2b2dfe..40bc2253b7f 100644 --- a/cookbook/91_tools/mcp/README.md +++ b/cookbook/91_tools/mcp/README.md @@ -111,4 +111,4 @@ You can modify these examples to: ## More Information - Read more about [MCP](https://modelcontextprotocol.io/introduction) -- Read about [Agno's MCP integration](https://docs.agno.com/mcp) +- Read about [Agno's MCP integration](https://docs.agno.com/tools/mcp/overview) diff --git a/cookbook/91_tools/mcp/mcp_toolbox_demo/requirements.txt b/cookbook/91_tools/mcp/mcp_toolbox_demo/requirements.txt index d967472291c..ab7ab3fe095 100644 --- a/cookbook/91_tools/mcp/mcp_toolbox_demo/requirements.txt +++ b/cookbook/91_tools/mcp/mcp_toolbox_demo/requirements.txt @@ -101,7 +101,7 @@ referencing==0.36.2 # via # jsonschema # jsonschema-specifications -requests==2.32.4 +requests==2.33.0 # via toolbox-core rpds-py==0.26.0 # via @@ -135,7 +135,7 @@ typing-inspection==0.4.1 # via # pydantic # pydantic-settings -urllib3==2.5.0 +urllib3==2.7.0 # via requests uvicorn==0.35.0 # via mcp diff --git a/cookbook/91_tools/mcp/mcp_toolbox_demo/uv.lock b/cookbook/91_tools/mcp/mcp_toolbox_demo/uv.lock index d11595323dc..c836635113a 100644 --- a/cookbook/91_tools/mcp/mcp_toolbox_demo/uv.lock +++ b/cookbook/91_tools/mcp/mcp_toolbox_demo/uv.lock @@ -999,7 +999,7 @@ wheels = [ [[package]] name = "requests" -version = "2.32.4" +version = "2.33.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "certifi" }, @@ -1007,9 +1007,9 @@ dependencies = [ { name = "idna" }, { name = "urllib3" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/e1/0a/929373653770d8a0d7ea76c37de6e41f11eb07559b103b1c02cafb3f7cf8/requests-2.32.4.tar.gz", hash = "sha256:27d0316682c8a29834d3264820024b62a36942083d52caf2f14c0591336d3422", size = 135258, upload-time = "2025-06-09T16:43:07.34Z" } +sdist = { url = "https://files.pythonhosted.org/packages/34/64/8860370b167a9721e8956ae116825caff829224fbca0ca6e7bf8ddef8430/requests-2.33.0.tar.gz", hash = "sha256:c7ebc5e8b0f21837386ad0e1c8fe8b829fa5f544d8df3b2253bff14ef29d7652", size = 134232, upload-time = "2026-03-25T15:10:41.586Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/7c/e4/56027c4a6b4ae70ca9de302488c5ca95ad4a39e190093d6c1a8ace08341b/requests-2.32.4-py3-none-any.whl", hash = "sha256:27babd3cda2a6d50b30443204ee89830707d396671944c998b5975b031ac2b2c", size = 64847, upload-time = "2025-06-09T16:43:05.728Z" }, + { url = "https://files.pythonhosted.org/packages/56/5d/c814546c2333ceea4ba42262d8c4d55763003e767fa169adc693bd524478/requests-2.33.0-py3-none-any.whl", hash = "sha256:3324635456fa185245e24865e810cecec7b4caf933d7eb133dcde67d48cee69b", size = 65017, upload-time = "2026-03-25T15:10:40.382Z" }, ] [[package]] @@ -1237,11 +1237,11 @@ wheels = [ [[package]] name = "urllib3" -version = "2.5.0" +version = "2.7.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/15/22/9ee70a2574a4f4599c47dd506532914ce044817c7752a79b6a51286319bc/urllib3-2.5.0.tar.gz", hash = "sha256:3fc47733c7e419d4bc3f6b3dc2b4f890bb743906a30d56ba4a5bfa4bbff92760", size = 393185, upload-time = "2025-06-18T14:07:41.644Z" } +sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602, upload-time = "2026-05-07T16:13:18.596Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a7/c2/fe1e52489ae3122415c51f387e221dd0773709bad6c6cdaa599e8a2c5185/urllib3-2.5.0-py3-none-any.whl", hash = "sha256:e6b01673c0fa6a13e374b50871808eb3bf7046c4b125b216f6bf1cc604cff0dc", size = 129795, upload-time = "2025-06-18T14:07:40.39Z" }, + { url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087, upload-time = "2026-05-07T16:13:17.151Z" }, ] [[package]] diff --git a/cookbook/91_tools/mcp/pipedream_google_calendar.py b/cookbook/91_tools/mcp/pipedream_google_calendar.py index fecc338055b..33e9e36a7aa 100644 --- a/cookbook/91_tools/mcp/pipedream_google_calendar.py +++ b/cookbook/91_tools/mcp/pipedream_google_calendar.py @@ -3,8 +3,8 @@ This example shows how to use Pipedream MCP servers (in this case the Google Calendar one) with Agno Agents. -1. Connect your Pipedream and Google Calendar accounts: https://mcp.pipedream.com/app/google_calendar -2. Get your Pipedream MCP server url: https://mcp.pipedream.com/app/google_calendar +1. Connect your Pipedream and Google Calendar accounts: https://pipedream.com/apps/google_calendar +2. Get your Pipedream MCP server url: https://pipedream.com/apps/google_calendar 3. Set the MCP_SERVER_URL environment variable to the MCP server url you got above 4. Install dependencies: uv pip install agno mcp """ diff --git a/cookbook/91_tools/mcp/pipedream_linkedin.py b/cookbook/91_tools/mcp/pipedream_linkedin.py index 56e0835a8f9..7af48bc2e6c 100644 --- a/cookbook/91_tools/mcp/pipedream_linkedin.py +++ b/cookbook/91_tools/mcp/pipedream_linkedin.py @@ -3,8 +3,8 @@ This example shows how to use Pipedream MCP servers (in this case the LinkedIn one) with Agno Agents. -1. Connect your Pipedream and LinkedIn accounts: https://mcp.pipedream.com/app/linkedin -2. Get your Pipedream MCP server url: https://mcp.pipedream.com/app/linkedin +1. Connect your Pipedream and LinkedIn accounts: https://pipedream.com/apps/linkedin +2. Get your Pipedream MCP server url: https://pipedream.com/apps/linkedin 3. Set the MCP_SERVER_URL environment variable to the MCP server url you got above 4. Install dependencies: uv pip install agno mcp """ diff --git a/cookbook/91_tools/mcp/pipedream_slack.py b/cookbook/91_tools/mcp/pipedream_slack.py index d76366c4538..d230a46172c 100644 --- a/cookbook/91_tools/mcp/pipedream_slack.py +++ b/cookbook/91_tools/mcp/pipedream_slack.py @@ -3,8 +3,8 @@ This example shows how to use Pipedream MCP servers (in this case the Slack one) with Agno Agents. -1. Connect your Pipedream and Slack accounts: https://mcp.pipedream.com/app/slack -2. Get your Pipedream MCP server url: https://mcp.pipedream.com/app/slack +1. Connect your Pipedream and Slack accounts: https://pipedream.com/apps/slack +2. Get your Pipedream MCP server url: https://pipedream.com/apps/slack 3. Set the MCP_SERVER_URL environment variable to the MCP server url you got above 4. Install dependencies: uv pip install agno mcp diff --git a/cookbook/91_tools/models/README.md b/cookbook/91_tools/models/README.md index 2b31b7962d1..e5733e16eb0 100644 --- a/cookbook/91_tools/models/README.md +++ b/cookbook/91_tools/models/README.md @@ -1,3 +1,13 @@ # models Cookbook examples for this tools subsection. + +| Example | Toolkit | What it shows | +|:--------|:--------|:--------------| +| `aimlapi_tools.py` | `AIMLAPITools` | Image, speech and transcription on AI/ML API through one agent; a second agent that makes a short video. Set `AIMLAPI_API_KEY`. | +| `azure_openai_tools.py` | `AzureOpenAITools` | Image generation on Azure OpenAI. | +| `gemini_image_generation.py` | `GeminiTools` | Image generation with Imagen. | +| `gemini_video_generation.py` | `GeminiTools` | Video generation with Veo (Vertex AI). | +| `morph.py` | `MorphTools` | Fast code edits with Morph. | +| `nebius_tools.py` | `NebiusTools` | Image generation on Nebius Token Factory. | +| `openai_tools.py` | `OpenAITools` | Transcription, image generation and speech with OpenAI. | diff --git a/cookbook/91_tools/models/TEST_LOG.md b/cookbook/91_tools/models/TEST_LOG.md index 0d339a2dfbe..20a954159fb 100644 --- a/cookbook/91_tools/models/TEST_LOG.md +++ b/cookbook/91_tools/models/TEST_LOG.md @@ -1,11 +1,38 @@ # Test Log +## Latest Verification — 2026-09-21 + +**Environment:** `.venv/bin/python` (Python 3.13), editable `libs/agno` 3.0.10 on branch `feat/aimlapi-tools` + +**Model:** `gpt-5.6-luna` via `AIMLAPI` + +**Key:** `AIMLAPI_API_KEY` (live gateway, `https://api.aimlapi.com`) + +--- + +### aimlapi_tools.py + +**Status:** PASS + +**Description:** One agent with `AIMLAPITools` (image + speech + transcription, `base_dir=tmp`), a second agent with only the video tool. Examples 1–3 run through `agent.run`, example 4 through `agenerate_video` directly. + +**Result:** +- Example 1 — `generate_image` called once; `image/png`, format `png`, 1.9 MB written to `tmp/`. +- Example 2 — `generate_speech` called once; `audio/mpeg`, format `mp3`, 44 KB. +- Example 3 — `transcribe_audio` on the file from example 2 (path relative to `base_dir`); transcript `the quick brown fox jumps over the lazy dog`, exact. +- Path guard — asking for `../../.zshrc` answers `Failed to transcribe audio: ... is outside the allowed directory tmp`; nothing is uploaded. +- Example 4 — `agenerate_video` (`bytedance/seedance-2-5`, 480p, 4 s): `video/mp4`, format `mp4`, 2.0 MB after 179 s of polling. + +Unit tests: `pytest libs/agno/tests/unit/tools/models/test_aimlapi.py` — 31 passed. + +--- + ### Pending **Status:** NOT RUN -**Description:** Tests for this cookbook directory have not been executed yet in this workspace. +**Description:** The other examples in this directory have not been executed yet in this workspace. -**Result:** Add individual run results after executing examples. +**Result:** Add individual run results after executing them. --- diff --git a/cookbook/91_tools/models/aimlapi_tools.py b/cookbook/91_tools/models/aimlapi_tools.py new file mode 100644 index 00000000000..5f1747d16c4 --- /dev/null +++ b/cookbook/91_tools/models/aimlapi_tools.py @@ -0,0 +1,102 @@ +"""Run `uv pip install agno` to install dependencies. + +AIMLAPITools gives an agent image, video, speech and transcription models from +AI/ML API (https://aimlapi.com) behind one key. Each capability is a separate +tool with its own model, so an agent can be handed only the ones it needs. + +Set AIMLAPI_API_KEY, or pass api_key=... to the toolkit. + +Example prompts to try: +- "Generate an image of a lighthouse in a storm" +- "Read this sentence aloud: The quick brown fox jumps over the lazy dog" +- "Make a short video of a paper boat drifting on a pond" +""" + +import mimetypes +from pathlib import Path + +from agno.agent import Agent +from agno.models.aimlapi import AIMLAPI +from agno.tools.models.aimlapi import AIMLAPITools + +OUTPUT_DIR = Path("tmp") + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- + +# The chat model and the media tools both run on AI/ML API. +agent = Agent( + model=AIMLAPI(id="gpt-5.6-luna"), + tools=[ + AIMLAPITools( + image_model="openai/gpt-image-2", + speech_model="openai/tts-1", + speech_voice="alloy", + # Local files handed to transcribe_audio are read from here only. + base_dir=OUTPUT_DIR, + # Video takes minutes; leave it off unless the agent should make clips. + enable_generate_video=False, + ) + ], + # The chat model does not take audio or video back as input; the generated + # media still comes out on the run output. + send_media_to_model=False, + markdown=True, +) + + +def save(artifact, stem: str) -> Path: + """Write a generated artifact next to the others, named by its media type.""" + extension = ( + mimetypes.guess_extension(artifact.mime_type or "") or f".{artifact.format}" + ) + path = OUTPUT_DIR / f"{stem}_{artifact.id}{extension}" + path.write_bytes(artifact.content) + return path + + +# --------------------------------------------------------------------------- +# Run Agent +# --------------------------------------------------------------------------- +if __name__ == "__main__": + OUTPUT_DIR.mkdir(exist_ok=True) + + # Example 1: image + response = agent.run("Generate an image of a lighthouse in a storm") + for image in response.images or []: + print(f"Image saved to {save(image, 'aimlapi')}") + + # Example 2: speech + response = agent.run( + "Read this aloud: The quick brown fox jumps over the lazy dog." + ) + saved_audio = [save(audio, "aimlapi") for audio in response.audio or []] + for path in saved_audio: + print(f"Audio saved to {path}") + + # Example 3: transcription of the speech we just made (path relative to base_dir) + for path in saved_audio: + agent.print_response(f"Transcribe the audio file {path.name}") + + # Example 4: video, on an agent that has the tool enabled + video_agent = Agent( + model=AIMLAPI(id="gpt-5.6-luna"), + tools=[ + AIMLAPITools( + enable_generate_image=False, + enable_generate_speech=False, + enable_transcribe_audio=False, + video_model="bytedance/seedance-2-5", + video_duration=4, + video_resolution="480p", + ) + ], + send_media_to_model=False, + markdown=True, + ) + response = video_agent.run( + "Make a short video of a paper boat drifting on a calm pond" + ) + for video in response.videos or []: + print(f"Video saved to {save(video, 'aimlapi')}") diff --git a/cookbook/observability/README.md b/cookbook/observability/README.md index 7d5227b0a0a..d9397eb9318 100644 --- a/cookbook/observability/README.md +++ b/cookbook/observability/README.md @@ -9,6 +9,7 @@ Observability examples for tracing and monitoring Agno agents, teams, and workfl - `arize_phoenix_via_openinference.py` - `arize_phoenix_via_openinference_local.py` - `atla_op.py` +- `confident_ai.py` - `langfuse_via_openinference.py` - `langfuse_via_openinference_response_model.py` - `langfuse_via_openlit.py` diff --git a/cookbook/observability/TEST_LOG.md b/cookbook/observability/TEST_LOG.md index d0ba9824291..9e9762c54d0 100644 --- a/cookbook/observability/TEST_LOG.md +++ b/cookbook/observability/TEST_LOG.md @@ -7,3 +7,12 @@ **Result:** Validation passed with zero violations. Runtime execution of individual cookbook scripts was not performed in this pass. --- +### confident_ai.py + +**Status:** PASS + +**Description:** Sends Agno agent traces to Confident AI through `confident-trace`. Calls `init()` once at startup, runs a HackerNews agent twice (one plain run, one inside `trace_context` with tags, metadata, and a user ID), and calls `shutdown()` in a `finally` block to flush spans. Verified with `ruff format`, `ruff check`, `cookbook/scripts/check_cookbook_pattern.py`, and a module import against `confident-trace==0.1.3` in `.venvs/demo`. + +**Result:** Static checks and import pass. Ran end to end with `CONFIDENT_API_KEY` and `OPENAI_API_KEY` set: both agent runs completed, tool and model calls executed, and span batches were accepted by the Confident AI collector with HTTP 200. Note that `CONFIDENT_OTEL_ENDPOINT` must match the project's region; a US project key against the EU endpoint returns 401 on export. + +--- diff --git a/cookbook/observability/confident_ai.py b/cookbook/observability/confident_ai.py new file mode 100644 index 00000000000..b4fc99d943a --- /dev/null +++ b/cookbook/observability/confident_ai.py @@ -0,0 +1,65 @@ +""" +Confident AI Observability Integration +====================================== + +Demonstrates sending Agno agent traces to Confident AI with confident-trace. + +confident-trace is an OpenTelemetry-native tracing SDK that detects Agno +automatically. Call init() once at startup and your agent, tool, and model +calls show up in the Confident AI Observatory with no other changes. + +Setup: + pip install agno openai confident-trace + +Set CONFIDENT_API_KEY and OPENAI_API_KEY before running this example. +For the EU region, also set CONFIDENT_OTEL_ENDPOINT to +https://eu.otel.confident-ai.com/v1/traces. +See https://www.confident-ai.com/docs/llm-tracing/introduction for details. +""" + +from agno.agent import Agent +from agno.models.openai import OpenAIResponses +from agno.tools.hackernews import HackerNewsTools +from confident_trace import init, shutdown, trace_context + +# --------------------------------------------------------------------------- +# Setup +# --------------------------------------------------------------------------- +# Initialize once at startup. This reads CONFIDENT_API_KEY from the +# environment and instruments Agno and the OpenAI SDK. +init() + + +# --------------------------------------------------------------------------- +# Create Agent +# --------------------------------------------------------------------------- +agent = Agent( + name="Hacker News Agent", + model=OpenAIResponses(id="gpt-5.6-luna"), + tools=[HackerNewsTools()], + instructions="You summarize Hacker News stories. Be concise and cite story titles.", + markdown=True, +) + + +# --------------------------------------------------------------------------- +# Run Example +# --------------------------------------------------------------------------- +if __name__ == "__main__": + try: + # A plain run is traced automatically. + agent.print_response("What are the top 3 stories on Hacker News right now?") + + # Attach tags, metadata, and a user ID to the trace before the run starts. + with trace_context( + tags=["cookbook"], + metadata={"release": "2026-09"}, + user_id="user-42", + ): + agent.print_response( + "Pick one of those stories and explain why it is trending.", + stream=True, + ) + finally: + # Flush pending spans before the process exits. + shutdown() diff --git a/libs/agno/agno/db/firestore/firestore.py b/libs/agno/agno/db/firestore/firestore.py index f6b641d607b..861586fd14a 100644 --- a/libs/agno/agno/db/firestore/firestore.py +++ b/libs/agno/agno/db/firestore/firestore.py @@ -765,7 +765,7 @@ def get_sessions( page (Optional[int]): The page number to get. sort_by (Optional[str]): The field to sort the sessions by. sort_order (Optional[str]): The order to sort the sessions by. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[AgentSession], List[TeamSession], List[WorkflowSession], Tuple[List[Dict[str, Any]], int]]: @@ -876,7 +876,7 @@ def rename_session( session_type (SessionType): The type of session to rename. session_name (str): The new name of the session. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Union[Session, Dict[str, Any]]]: @@ -1244,7 +1244,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user (optional, for filtering). Returns: @@ -1305,7 +1305,7 @@ def get_user_memories( page (Optional[int]): The page number to get. sort_by (Optional[str]): The field to sort the memories by. sort_order (Optional[str]): The order to sort the memories by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. create_table_if_not_found: Whether to create the index if it doesn't exist. Returns: @@ -1429,7 +1429,7 @@ def upsert_user_memory( Args: memory (UserMemory): The memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -1991,7 +1991,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2058,7 +2058,7 @@ def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type of eval to filter by. filter_type (Optional[EvalFilterType]): The type of filter to apply. - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: @@ -2146,7 +2146,7 @@ def rename_eval_run( Args: eval_run_id (str): The ID of the eval run to update. name (str): The new name of the eval run. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only rename the run if owned by this user. Returns: diff --git a/libs/agno/agno/db/migrations/V3_MIGRATION_GUIDE.md b/libs/agno/agno/db/migrations/V3_MIGRATION_GUIDE.md index 91ae99fc18c..54c5b53cd77 100644 --- a/libs/agno/agno/db/migrations/V3_MIGRATION_GUIDE.md +++ b/libs/agno/agno/db/migrations/V3_MIGRATION_GUIDE.md @@ -197,7 +197,7 @@ db.cleanup_legacy_runs_column() ``` This refuses to drop the column if any session still has non-null legacy `runs` -content (a sign that that session was not migrated). If you really want to force +content (a sign that the session was not migrated). If you really want to force it anyway: ```python diff --git a/libs/agno/agno/db/mongo/async_mongo.py b/libs/agno/agno/db/mongo/async_mongo.py index 3014cd90971..a61c2343458 100644 --- a/libs/agno/agno/db/mongo/async_mongo.py +++ b/libs/agno/agno/db/mongo/async_mongo.py @@ -992,7 +992,7 @@ async def get_sessions( page (Optional[int]): The page number to get. sort_by (Optional[str]): The field to sort the sessions by. sort_order (Optional[str]): The order to sort the sessions by. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[AgentSession], List[TeamSession], List[WorkflowSession], Tuple[List[Dict[str, Any]], int]]: @@ -1516,7 +1516,7 @@ async def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user to verify ownership. If provided, only return the memory if it belongs to this user. Returns: @@ -1573,7 +1573,7 @@ async def get_user_memories( page (Optional[int]): The page number to get. sort_by (Optional[str]): The field to sort the memories by. sort_order (Optional[str]): The order to sort the memories by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: Union[List[UserMemory], Tuple[List[Dict[str, Any]], int]]: @@ -1703,7 +1703,7 @@ async def upsert_user_memory( Args: memory (UserMemory): The memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2313,7 +2313,7 @@ async def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2375,7 +2375,7 @@ async def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type of eval to filter by. filter_type (Optional[EvalFilterType]): The type of filter to apply. - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. Returns: Union[List[EvalRunRecord], Tuple[List[Dict[str, Any]], int]]: @@ -2452,7 +2452,7 @@ async def rename_eval_run( Args: eval_run_id (str): The ID of the eval run to update. name (str): The new name of the eval run. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only rename the run if owned by this user. Returns: diff --git a/libs/agno/agno/db/mongo/mongo.py b/libs/agno/agno/db/mongo/mongo.py index c80284aec25..ae37dc3f425 100644 --- a/libs/agno/agno/db/mongo/mongo.py +++ b/libs/agno/agno/db/mongo/mongo.py @@ -802,7 +802,7 @@ def get_sessions( page (Optional[int]): The page number to get. sort_by (Optional[str]): The field to sort the sessions by. sort_order (Optional[str]): The order to sort the sessions by. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the collection if it doesn't exist. Returns: @@ -1323,7 +1323,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user to verify ownership. If provided, only return the memory if it belongs to this user. Returns: @@ -1380,7 +1380,7 @@ def get_user_memories( page (Optional[int]): The page number to get. sort_by (Optional[str]): The field to sort the memories by. sort_order (Optional[str]): The order to sort the memories by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. create_table_if_not_found: Whether to create the collection if it doesn't exist. Returns: @@ -1509,7 +1509,7 @@ def upsert_user_memory( Args: memory (UserMemory): The memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2118,7 +2118,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2180,7 +2180,7 @@ def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type of eval to filter by. filter_type (Optional[EvalFilterType]): The type of filter to apply. - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the collection if it doesn't exist. Returns: @@ -2258,7 +2258,7 @@ def rename_eval_run( Args: eval_run_id (str): The ID of the eval run to update. name (str): The new name of the eval run. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only rename the run if owned by this user. Returns: diff --git a/libs/agno/agno/db/mysql/async_mysql.py b/libs/agno/agno/db/mysql/async_mysql.py index 50fc5a99bd5..11a6341db06 100644 --- a/libs/agno/agno/db/mysql/async_mysql.py +++ b/libs/agno/agno/db/mysql/async_mysql.py @@ -949,7 +949,7 @@ async def get_session( session_id (str): ID of the session to read. session_type (Optional[SessionType]): Type of session to get. Defaults to None. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Union[Session, Dict[str, Any], None]: @@ -1055,7 +1055,7 @@ async def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[Session], Tuple[List[Dict], int]]: @@ -1158,7 +1158,7 @@ async def rename_session( session_type (SessionType): The type of session to rename. session_name (str): The new name for the session. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Union[Session, Dict[str, Any]]]: @@ -1651,7 +1651,7 @@ async def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Union[UserMemory, Dict[str, Any], None]: @@ -1711,7 +1711,7 @@ async def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: Union[List[UserMemory], Tuple[List[Dict[str, Any]], int]]: @@ -1869,7 +1869,7 @@ async def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2567,7 +2567,7 @@ async def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2631,7 +2631,7 @@ async def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. Returns: Union[List[EvalRunRecord], Tuple[List[Dict[str, Any]], int]]: diff --git a/libs/agno/agno/db/mysql/mysql.py b/libs/agno/agno/db/mysql/mysql.py index ea12048cb44..d82d0a7479e 100644 --- a/libs/agno/agno/db/mysql/mysql.py +++ b/libs/agno/agno/db/mysql/mysql.py @@ -946,7 +946,7 @@ def get_session( session_id (str): ID of the session to read. session_type (Optional[SessionType]): Type of session to get. Defaults to None. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Union[Session, Dict[str, Any], None]: @@ -1050,7 +1050,7 @@ def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[Session], Tuple[List[Dict], int]]: @@ -1152,7 +1152,7 @@ def rename_session( session_type (SessionType): The type of session to rename. session_name (str): The new name for the session. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Union[Session, Dict[str, Any]]]: @@ -1633,7 +1633,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The user ID to filter by. Defaults to None. Returns: @@ -1692,7 +1692,7 @@ def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: @@ -1844,7 +1844,7 @@ def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2537,7 +2537,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2600,7 +2600,7 @@ def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: diff --git a/libs/agno/agno/db/postgres/async_postgres.py b/libs/agno/agno/db/postgres/async_postgres.py index 98af5f40933..4bda764f110 100644 --- a/libs/agno/agno/db/postgres/async_postgres.py +++ b/libs/agno/agno/db/postgres/async_postgres.py @@ -1300,7 +1300,7 @@ async def get_session( session_id (str): ID of the session to read. user_id (Optional[str]): User ID to filter by. Defaults to None. session_type (Optional[SessionType]): Type of session to read. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. runs_limit (Optional[int]): If set, attach only the most recent ``runs_limit`` runs instead of the full history. For a fully-migrated session this is an indexed ``ORDER BY run_index DESC LIMIT`` query; for a session that still @@ -1418,7 +1418,7 @@ async def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[Session], Tuple[List[Dict], int]]: @@ -1826,7 +1826,7 @@ async def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user to filter by. Returns: @@ -1887,7 +1887,7 @@ async def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: Union[List[UserMemory], Tuple[List[Dict[str, Any]], int]]: @@ -2042,7 +2042,7 @@ async def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2690,7 +2690,7 @@ async def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2753,7 +2753,7 @@ async def get_eval_runs( model_id (Optional[str]): The ID of the model to filter by. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. user_id (Optional[str]): If set, only return runs owned by this user. Returns: diff --git a/libs/agno/agno/db/postgres/postgres.py b/libs/agno/agno/db/postgres/postgres.py index de1744bbacb..2804c95041b 100644 --- a/libs/agno/agno/db/postgres/postgres.py +++ b/libs/agno/agno/db/postgres/postgres.py @@ -1493,7 +1493,7 @@ def get_session( session_id (str): ID of the session to read. session_type (SessionType): Type of session to get. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. runs_limit (Optional[int]): If set, attach only the most recent ``runs_limit`` runs instead of the full history. For a fully-migrated session this is an indexed ``ORDER BY run_index DESC LIMIT`` query; for a session that still @@ -1609,7 +1609,7 @@ def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[Session], Tuple[List[Dict], int]]: @@ -1714,7 +1714,7 @@ def rename_session( session_type (Optional[SessionType]): The type of session to rename. Defaults to None. session_name (str): The new name for the session. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Union[Session, Dict[str, Any]]]: @@ -2241,7 +2241,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user to filter by. Defaults to None. Returns: @@ -2302,7 +2302,7 @@ def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: @@ -2456,7 +2456,7 @@ def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -3193,7 +3193,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -3255,7 +3255,7 @@ def get_eval_runs( model_id (Optional[str]): The ID of the model to filter by. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. user_id (Optional[str]): If set, only return runs owned by this user. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. diff --git a/libs/agno/agno/db/singlestore/singlestore.py b/libs/agno/agno/db/singlestore/singlestore.py index bdc2aa0e9f5..f3da76f9e1f 100644 --- a/libs/agno/agno/db/singlestore/singlestore.py +++ b/libs/agno/agno/db/singlestore/singlestore.py @@ -1001,7 +1001,7 @@ def get_session( session_id (str): ID of the session to read. session_type (Optional[SessionType]): Type of session to get. If None, the type is inferred. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Union[Session, Dict[str, Any], None]: @@ -1105,7 +1105,7 @@ def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: @@ -1208,7 +1208,7 @@ def rename_session( session_type (SessionType): The type of session to rename. session_name (str): The new name for the session. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Union[Session, Dict[str, Any]]]: @@ -1672,7 +1672,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user to filter by. Defaults to None. Returns: @@ -1731,7 +1731,7 @@ def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: @@ -1868,7 +1868,7 @@ def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2533,7 +2533,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2596,7 +2596,7 @@ def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: diff --git a/libs/agno/agno/db/sqlite/async_sqlite.py b/libs/agno/agno/db/sqlite/async_sqlite.py index 614ae876c37..e3ef778cd97 100644 --- a/libs/agno/agno/db/sqlite/async_sqlite.py +++ b/libs/agno/agno/db/sqlite/async_sqlite.py @@ -1295,7 +1295,7 @@ async def get_session( session_id (str): ID of the session to read. session_type (SessionType): Type of session to get. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. runs_limit (Optional[int]): If set, attach only the most recent ``runs_limit`` runs instead of the full history. For a fully-migrated session this is an indexed ``ORDER BY run_index DESC LIMIT`` query; for a session that still @@ -1417,7 +1417,7 @@ async def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: List[Session]: @@ -1561,7 +1561,7 @@ async def upsert_session( Args: session (Session): The session data to upsert. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Session]: @@ -2017,7 +2017,7 @@ async def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The user ID to filter by. Defaults to None. Returns: @@ -2076,7 +2076,7 @@ async def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: @@ -2217,7 +2217,7 @@ async def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -2936,7 +2936,7 @@ async def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -2999,7 +2999,7 @@ async def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: @@ -3077,7 +3077,7 @@ async def rename_eval_run( Args: eval_run_id (str): The ID of the eval run to update. name (str): The new name of the eval run. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only rename the run if owned by this user. Returns: diff --git a/libs/agno/agno/db/sqlite/sqlite.py b/libs/agno/agno/db/sqlite/sqlite.py index 8d04d96f84a..9c032d733ee 100644 --- a/libs/agno/agno/db/sqlite/sqlite.py +++ b/libs/agno/agno/db/sqlite/sqlite.py @@ -1508,7 +1508,7 @@ def get_session( session_id (str): ID of the session to read. session_type (SessionType): Type of session to get. user_id (Optional[str]): User ID to filter by. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. runs_limit (Optional[int]): If set, attach only the most recent ``runs_limit`` runs instead of the full history. For a fully-migrated session this is an indexed ``ORDER BY run_index DESC LIMIT`` query; for a session that still @@ -1625,7 +1625,7 @@ def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: @@ -1769,7 +1769,7 @@ def upsert_session( Args: session (Session): The session data to upsert. - deserialize (Optional[bool]): Whether to serialize the session. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the session. Defaults to True. Returns: Optional[Session]: @@ -2225,7 +2225,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The user ID to filter by. Defaults to None. Returns: @@ -2284,7 +2284,7 @@ def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: @@ -2424,7 +2424,7 @@ def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -3140,7 +3140,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -3203,7 +3203,7 @@ def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type(s) of eval to filter by. filter_type (Optional[EvalFilterType]): Filter by component type (agent, team, workflow). - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. create_table_if_not_found (Optional[bool]): Whether to create the table if it doesn't exist. Returns: @@ -3281,7 +3281,7 @@ def rename_eval_run( Args: eval_run_id (str): The ID of the eval run to update. name (str): The new name of the eval run. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only rename the run if owned by this user. Returns: diff --git a/libs/agno/agno/db/surrealdb/surrealdb.py b/libs/agno/agno/db/surrealdb/surrealdb.py index 1d657be67ba..d8b3b46ca27 100644 --- a/libs/agno/agno/db/surrealdb/surrealdb.py +++ b/libs/agno/agno/db/surrealdb/surrealdb.py @@ -655,7 +655,7 @@ def get_sessions( page (Optional[int]): The page number to return. Defaults to None. sort_by (Optional[str]): The field to sort by. Defaults to None. sort_order (Optional[str]): The sort order. Defaults to None. - deserialize (Optional[bool]): Whether to serialize the sessions. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the sessions. Defaults to True. Returns: Union[List[Session], Tuple[List[Dict], int]]: @@ -1035,7 +1035,7 @@ def get_user_memory( Args: memory_id (str): The ID of the memory to get. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. user_id (Optional[str]): The ID of the user to filter by. Defaults to None. Returns: @@ -1086,7 +1086,7 @@ def get_user_memories( page (Optional[int]): The page number. sort_by (Optional[str]): The column to sort by. sort_order (Optional[str]): The order to sort by. - deserialize (Optional[bool]): Whether to serialize the memories. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memories. Defaults to True. Returns: @@ -1205,7 +1205,7 @@ def upsert_user_memory( Args: memory (UserMemory): The user memory to upsert. - deserialize (Optional[bool]): Whether to serialize the memory. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the memory. Defaults to True. Returns: Optional[Union[UserMemory, Dict[str, Any]]]: @@ -1607,7 +1607,7 @@ def get_eval_run( Args: eval_run_id (str): The ID of the eval run to get. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only return the run if owned by this user. Returns: @@ -1656,7 +1656,7 @@ def get_eval_runs( user_id (Optional[str]): If set, only return runs owned by this user. eval_type (Optional[List[EvalType]]): The type of eval to filter by. filter_type (Optional[EvalFilterType]): The type of filter to apply. - deserialize (Optional[bool]): Whether to serialize the eval runs. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval runs. Defaults to True. Returns: Union[List[EvalRunRecord], Tuple[List[Dict[str, Any]], int]]: @@ -1717,7 +1717,7 @@ def rename_eval_run( Args: eval_run_id (str): The ID of the eval run to update. name (str): The new name of the eval run. - deserialize (Optional[bool]): Whether to serialize the eval run. Defaults to True. + deserialize (Optional[bool]): Whether to deserialize the eval run. Defaults to True. user_id (Optional[str]): If set, only rename the run if owned by this user. Returns: diff --git a/libs/agno/agno/knowledge/chunking/row.py b/libs/agno/agno/knowledge/chunking/row.py index 8e3e4a7bb0a..3b97118ea5b 100644 --- a/libs/agno/agno/knowledge/chunking/row.py +++ b/libs/agno/agno/knowledge/chunking/row.py @@ -17,12 +17,11 @@ def chunk(self, document: Document) -> List[Document]: raise ValueError("Document content must be a string") rows = document.content.splitlines() + start_index = document.meta_data.get("start_row", 1) # Set by readers that split a file into pages - if self.skip_header and rows: + if self.skip_header and rows and start_index == 1: # Only a document starting at row 1 holds the header rows = rows[1:] - start_index = 2 - else: - start_index = 1 + start_index += 1 chunks = [] for i, row in enumerate(rows): diff --git a/libs/agno/agno/knowledge/chunking/strategy.py b/libs/agno/agno/knowledge/chunking/strategy.py index dbb2aea5d63..02a25f1ba19 100644 --- a/libs/agno/agno/knowledge/chunking/strategy.py +++ b/libs/agno/agno/knowledge/chunking/strategy.py @@ -36,13 +36,13 @@ async def achunk(self, document: Document) -> List[Document]: return self.chunk(document) def clean_text(self, text: str) -> str: - """Clean the text by replacing multiple newlines with a single newline""" + """Clean the text by collapsing runs of each whitespace character type.""" import re # Replace multiple newlines with a single newline cleaned_text = re.sub(r"\n+", "\n", text) # Replace multiple spaces with a single space - cleaned_text = re.sub(r"\s+", " ", cleaned_text) + cleaned_text = re.sub(r" +", " ", cleaned_text) # Replace multiple tabs with a single tab cleaned_text = re.sub(r"\t+", "\t", cleaned_text) # Replace multiple carriage returns with a single carriage return diff --git a/libs/agno/agno/knowledge/knowledge.py b/libs/agno/agno/knowledge/knowledge.py index 23648e3fa5f..da6af6300bd 100644 --- a/libs/agno/agno/knowledge/knowledge.py +++ b/libs/agno/agno/knowledge/knowledge.py @@ -4,6 +4,7 @@ import json import math import time +from contextlib import contextmanager from dataclasses import dataclass from enum import Enum from io import BytesIO @@ -27,6 +28,7 @@ RemoteContent, ) from agno.knowledge.remote_knowledge import RemoteKnowledge +from agno.knowledge.reranker.base import Reranker from agno.knowledge.types import ContentType from agno.knowledge.utils import get_agno_metadata, merge_user_metadata, set_agno_metadata, strip_agno_metadata from agno.utils.http import async_fetch_with_retry @@ -35,6 +37,8 @@ from agno.utils.string import generate_id ContentDict = Dict[str, Union[str, Dict[str, str]]] +# PageCoordinator.search rejects a limit outside 1..20. +_MAX_PAGE_SEARCH_LIMIT = 20 _DATABASE_UNSET = object() @@ -79,6 +83,11 @@ class Knowledge(RemoteKnowledge): page_store: Optional[Any] = None page_search: Optional[PageSearchConfig] = None + # Reorders results after the vector db returns them, so a strategy that needs to + # compare candidates against each other (diversity, recency) sees a real pool. + # Runs after any reranker configured on the vector db itself. + reranker: Optional[Reranker] = None + def __init__( self, *, @@ -94,6 +103,7 @@ def __init__( page_search: Optional[PageSearchConfig] = None, max_embedding_retries: int = 0, embedding_retry_backoff: float = 1.0, + reranker: Optional[Reranker] = None, contents_db: Optional[Union[BaseDb, AsyncBaseDb]] = cast(Any, _DATABASE_UNSET), ): """Configure Knowledge using keyword arguments. @@ -117,6 +127,14 @@ def __init__( self.embedding_retry_backoff = embedding_retry_backoff self.page_store = page_store self.page_search = page_search + self.reranker = reranker + if reranker is not None and getattr(vector_db, "reranker", None) is not None: + log_warning( + "A reranker is set on both Knowledge and the vector db. Only the one on " + "Knowledge is applied and the vector db's is ignored: running both would " + "rerank a pool that was already reordered. Prefer the one on Knowledge, " + "which works with every vector db and can widen the candidate pool." + ) self.__post_init__() @property @@ -165,6 +183,70 @@ def _page_documents(result: SearchResult) -> List[Document]: for hit in result.results ] + def _page_search_limit(self, max_results: int) -> int: + """Widen within the page search ceiling, which rejects a limit above 20. + + The gain is small either way: search_pages drops hits from the tail until the + serialized result fits its byte budget, so a page fetch is bounded well before + the ceiling. + """ + return min(self._search_limit(max_results), _MAX_PAGE_SEARCH_LIMIT) + + @contextmanager + def _vector_db_reranker_suspended(self): + """Skip the vector db's own reranker while the one on Knowledge is in charge. + + Knowledge widens the fetch for its reranker, so letting the vector db reorder + and trim that pool first would discard the candidates it was widened for. + """ + if self.reranker is None or getattr(self.vector_db, "reranker", None) is None: + yield + return + from agno.vectordb.base import suppress_reranker + + with suppress_reranker(): + yield + + def _search_limit(self, max_results: int) -> int: + """Widen the vector db fetch so the reranker has candidates to choose between.""" + if self.reranker is None: + return max_results + # The reranker decides how wide its own pool needs to be. + return self.reranker.search_limit(max_results) + + def _rerank_documents(self, query: str, documents: List[Document], max_results: int) -> List[Document]: + """Apply the knowledge-level reranker, then trim to the caller's requested count.""" + if self.reranker is None: + # Unchanged from before this hook existed: the adapter already applied the limit. + return documents + try: + kwargs = {"limit": max_results} if self.reranker.accepts_limit() else {} + reranked = self.reranker.rerank(query=query, documents=documents, **kwargs) + except ValueError: + # A misconfigured reranker would otherwise look like it ran and changed nothing. + raise + except Exception as e: + # A reranker failure degrades ordering, not availability: keep the vector db order. + log_error(f"Error reranking documents: {str(e)}") + return documents[:max_results] + return reranked[:max_results] + + async def _arerank_documents(self, query: str, documents: List[Document], max_results: int) -> List[Document]: + """Async variant of ``_rerank_documents``.""" + if self.reranker is None: + # See the matching comment in ``_rerank_documents``. + return documents + try: + # arerank always accepts limit; it forwards only to a rerank that takes it. + reranked = await self.reranker.arerank(query=query, documents=documents, limit=max_results) + except ValueError: + # See the matching comment in ``_rerank_documents``. + raise + except Exception as e: + log_error(f"Error reranking documents: {str(e)}") + return documents[:max_results] + return reranked[:max_results] + def setup(self) -> None: """Prepare and validate coordinated page storage before query traffic.""" pages = self._pages() @@ -948,9 +1030,9 @@ def search( if self.page_store is not None: if filters: raise ValueError("Page knowledge does not support filters") - return self._page_documents( - self.search_pages(query, limit=max_results if max_results is not None else self.max_results) - ) + page_limit = max_results if max_results is not None else self.max_results + page_documents = self._page_documents(self.search_pages(query, limit=self._page_search_limit(page_limit))) + return self._rerank_documents(query, page_documents, page_limit) from agno.vectordb import VectorDb from agno.vectordb.search import SearchType @@ -971,12 +1053,14 @@ def search( _max_results = max_results or self.max_results log_debug(f"Getting {_max_results} relevant documents for query: {query}") - return self.vector_db.search( - query=query, - limit=_max_results, - filters=search_filters, - **strict_user_id_kwarg(self.vector_db.search, user_id), - ) + with self._vector_db_reranker_suspended(): + documents = self.vector_db.search( + query=query, + limit=self._search_limit(_max_results), + filters=search_filters, + **strict_user_id_kwarg(self.vector_db.search, user_id), + ) + return self._rerank_documents(query, documents, _max_results) except ValueError: # The adapters raise these outside their own catch-alls on purpose. raise @@ -1000,9 +1084,11 @@ async def asearch( if self.page_store is not None: if filters: raise ValueError("Page knowledge does not support filters") - return self._page_documents( - await self.asearch_pages(query, limit=max_results if max_results is not None else self.max_results) + page_limit = max_results if max_results is not None else self.max_results + page_documents = self._page_documents( + await self.asearch_pages(query, limit=self._page_search_limit(page_limit)) ) + return await self._arerank_documents(query, page_documents, page_limit) from agno.vectordb import VectorDb from agno.vectordb.search import SearchType @@ -1022,21 +1108,24 @@ async def asearch( _max_results = max_results or self.max_results log_debug(f"Getting {_max_results} relevant documents for query: {query}") - try: - return await self.vector_db.async_search( - query=query, - limit=_max_results, - filters=search_filters, - **strict_user_id_kwarg(self.vector_db.async_search, user_id), - ) - except NotImplementedError: - log_info("Vector db does not support async search") - return self.vector_db.search( - query=query, - limit=_max_results, - filters=search_filters, - **strict_user_id_kwarg(self.vector_db.search, user_id), - ) + search_limit = self._search_limit(_max_results) + with self._vector_db_reranker_suspended(): + try: + documents = await self.vector_db.async_search( + query=query, + limit=search_limit, + filters=search_filters, + **strict_user_id_kwarg(self.vector_db.async_search, user_id), + ) + except NotImplementedError: + log_info("Vector db does not support async search") + documents = self.vector_db.search( + query=query, + limit=search_limit, + filters=search_filters, + **strict_user_id_kwarg(self.vector_db.search, user_id), + ) + return await self._arerank_documents(query, documents, _max_results) except ValueError: # See the matching comment in ``search``. raise diff --git a/libs/agno/agno/knowledge/reader/sitemap_reader.py b/libs/agno/agno/knowledge/reader/sitemap_reader.py index 1e5137043a8..38a6c28d611 100644 --- a/libs/agno/agno/knowledge/reader/sitemap_reader.py +++ b/libs/agno/agno/knowledge/reader/sitemap_reader.py @@ -11,6 +11,7 @@ """ import gzip +import zlib from typing import Generator, List, Optional, Tuple from urllib.parse import urlparse from xml.etree import ElementTree @@ -108,7 +109,7 @@ def _decode_sitemap_bytes(raw: bytes) -> bytes: if raw[:2] == _GZIP_MAGIC: try: return gzip.decompress(raw) - except OSError: + except (OSError, EOFError, zlib.error): return raw return raw diff --git a/libs/agno/agno/knowledge/reader/text_reader.py b/libs/agno/agno/knowledge/reader/text_reader.py index eee28a968f8..6e0a2b34a56 100644 --- a/libs/agno/agno/knowledge/reader/text_reader.py +++ b/libs/agno/agno/knowledge/reader/text_reader.py @@ -48,7 +48,9 @@ def read(self, file: Union[Path, IO[Any]], name: Optional[str] = None) -> List[D log_debug(f"Reading uploaded file: {getattr(file, 'name', 'BytesIO')}") file_name = name or getattr(file, "name", "text_file").split(".")[0] file.seek(0) - file_contents = file.read().decode(self.encoding or "utf-8") + file_contents = file.read() + if isinstance(file_contents, bytes): + file_contents = file_contents.decode(self.encoding or "utf-8") documents = [ Document( @@ -88,7 +90,9 @@ async def async_read(self, file: Union[Path, IO[Any]], name: Optional[str] = Non log_debug(f"Reading uploaded file asynchronously: {getattr(file, 'name', 'BytesIO')}") file_name = name or getattr(file, "name", "text_file").split(".")[0] file.seek(0) - file_contents = file.read().decode(self.encoding or "utf-8") + file_contents = file.read() + if isinstance(file_contents, bytes): + file_contents = file_contents.decode(self.encoding or "utf-8") document = Document( name=file_name, diff --git a/libs/agno/agno/knowledge/reranker/__init__.py b/libs/agno/agno/knowledge/reranker/__init__.py index fc94da5aa90..1967b0dc7f2 100644 --- a/libs/agno/agno/knowledge/reranker/__init__.py +++ b/libs/agno/agno/knowledge/reranker/__init__.py @@ -1,3 +1,4 @@ from agno.knowledge.reranker.base import Reranker +from agno.knowledge.reranker.mmr import MMRReranker -__all__ = ["Reranker"] +__all__ = ["Reranker", "MMRReranker"] diff --git a/libs/agno/agno/knowledge/reranker/base.py b/libs/agno/agno/knowledge/reranker/base.py index fa55f137094..a8a9cd590e8 100644 --- a/libs/agno/agno/knowledge/reranker/base.py +++ b/libs/agno/agno/knowledge/reranker/base.py @@ -1,6 +1,8 @@ -from typing import List +import asyncio +from inspect import signature +from typing import List, Optional -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from agno.knowledge.document import Document @@ -10,5 +12,39 @@ class Reranker(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True, populate_by_name=True) - def rerank(self, query: str, documents: List[Document]) -> List[Document]: + # Candidates fetched per requested result. Reranking is worth its cost because it + # rescues documents the vector search ranked below the cutoff, so the default asks + # for a wider pool; set it to 1 for a reranker that only needs to reorder. + candidate_multiplier: int = Field(default=3, ge=1) + # Ceiling on the widened fetch, so a large request cannot turn one search into an + # unbounded scan. + max_candidates: int = Field(default=100, ge=1) + + def search_limit(self, max_results: int) -> int: + """The number of candidates the vector db should return for this reranker.""" + # The ceiling caps the widening, never the caller's own request: clamping below + # max_results would return fewer documents than were asked for. + return max(min(max_results * self.candidate_multiplier, self.max_candidates), max_results) + + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + """Reorder documents. ``limit`` is the count the caller keeps, which a reranker + that selects a subset can use to stop early; scoring rerankers ignore it.""" raise NotImplementedError + + def accepts_limit(self) -> bool: + """Whether this reranker's ``rerank`` takes the caller's kept count. + + Rerankers written against the older two-argument signature, including ones + outside this repo, are still called without it. + """ + try: + return "limit" in signature(self.rerank).parameters + except (TypeError, ValueError): + return False + + async def arerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + """Async rerank. Runs the sync implementation off the event loop, since a + reranker that calls a provider would otherwise block it.""" + if self.accepts_limit(): + return await asyncio.to_thread(self.rerank, query, documents, limit) + return await asyncio.to_thread(self.rerank, query, documents) diff --git a/libs/agno/agno/knowledge/reranker/mmr.py b/libs/agno/agno/knowledge/reranker/mmr.py new file mode 100644 index 00000000000..69dfd7685c9 --- /dev/null +++ b/libs/agno/agno/knowledge/reranker/mmr.py @@ -0,0 +1,182 @@ +import asyncio +from dataclasses import replace +from typing import Any, List, Optional, Tuple + +from pydantic import Field, field_validator + +from agno.knowledge.document import Document +from agno.knowledge.reranker.base import Reranker +from agno.utils.vectors import dot, unit + + +class MMRReranker(Reranker): + """Selects results that are relevant to the query but unlike each other. + + Plain vector search returns the closest matches, which are often near-duplicates of + one another. MMR picks documents one at a time, discounting each candidate by how + similar it already is to what has been selected. + + Requires an embedding on every candidate document. Verified against live backends: + pgvector, Qdrant (vector and hybrid), Chroma and LanceDB return them; Milvus, + MongoDB, Redis, Valkey and Qdrant keyword search do not, and MMR raises there rather + than returning an unreranked list. Pinecone omits vectors unless the store is built + with return_vectors=True. + + It also needs an embedder to embed the query. Vector dbs that embed queries + themselves (Upstash hosted embeddings) expose none, so MMR cannot run there. + + The query is embedded with the embedder attached to the search results, so it always + uses the same model the documents were indexed with. + + Returned documents are shallow copies carrying the MMR score; meta_data and embedding + are shared with the inputs. + + Do not re-sort the result by reranking_score. Other rerankers score each document + independently, so their order can be rebuilt from the scores; MMR chooses each + document against the ones already chosen, so its scores are not descending and + sorting by them discards the diversity ordering. + """ + + # Selection compares candidates against each other, so it needs a pool wider than + # the caller asked for: a document can only be surfaced if it was retrieved. + candidate_multiplier: int = Field(default=5, ge=1) + + # Weight between relevance and diversity: 1.0 ranks by relevance alone, 0.0 by + # difference alone. + lambda_mult: float = Field(default=0.5, ge=0.0, le=1.0) + # Caps how many documents are selected. Leave unset on Knowledge.reranker, which + # trims to max_results anyway: a smaller top_n returns fewer documents than asked for. + top_n: Optional[int] = Field(default=None, gt=0) + + @field_validator("lambda_mult", mode="before") + @classmethod + def _reject_bool_lambda(cls, value: Any) -> Any: + # bool is an int subclass, so True would otherwise coerce to 1.0. + if isinstance(value, bool): + raise ValueError("lambda_mult must be a number between 0.0 and 1.0") + return value + + @field_validator("top_n", mode="before") + @classmethod + def _reject_bool_top_n(cls, value: Any) -> Any: + if isinstance(value, bool): + raise ValueError("top_n must be a positive integer") + return value + + def _select(self, query_embedding: Optional[List[float]], documents: List[Document], limit: int) -> List[Document]: + # Vector dbs return embeddings as lists or as numpy arrays, whose truth value + # is ambiguous, so length is the portable emptiness check throughout. + if query_embedding is None or len(query_embedding) == 0: + raise ValueError("MMRReranker could not embed the query: the embedder returned no vector") + + raw: List[List[float]] = [doc.embedding for doc in documents] # type: ignore[misc] + # zip() would silently truncate to the shorter vector and score against a prefix. + dimensions = {len(embedding) for embedding in raw} | {len(query_embedding)} + if len(dimensions) > 1: + raise ValueError( + f"MMRReranker requires embeddings of one dimension, but got {sorted(dimensions)}. " + "The query embedder and the indexed documents likely use different models." + ) + + # Normalise once: every similarity below is then a dot product, instead of + # recomputing the same norms across thousands of pair comparisons. + embeddings = [unit(embedding) for embedding in raw] + unit_query = unit(query_embedding) + relevance = [dot(unit_query, embedding) for embedding in embeddings] + + selected: List[Tuple[int, float]] = [] + remaining = list(range(len(documents))) + + # Seed with the closest match to the query. Scoring the first pick with the MMR + # formula would tie every candidate at lambda_mult=0.0 and pick by input order. + first = max(remaining, key=lambda candidate: relevance[candidate]) + selected.append((first, relevance[first])) + remaining.remove(first) + + # Each candidate's similarity to the nearest selected document, extended as + # documents are picked. Recomputing it per iteration is quadratic in the number + # selected, which at the candidate ceiling dominates the search itself. + best_redundancy = [0.0] * len(documents) + for candidate in remaining: + best_redundancy[candidate] = dot(embeddings[candidate], embeddings[first]) + + while remaining and len(selected) < limit: + best_index = remaining[0] + best_score = float("-inf") + for candidate in remaining: + score = self.lambda_mult * relevance[candidate] - (1.0 - self.lambda_mult) * best_redundancy[candidate] + if score > best_score: + best_score = score + best_index = candidate + selected.append((best_index, best_score)) + remaining.remove(best_index) + for candidate in remaining: + similarity = dot(embeddings[candidate], embeddings[best_index]) + if similarity > best_redundancy[candidate]: + best_redundancy[candidate] = similarity + + results: List[Document] = [] + for index, score in selected: + # A shallow copy, made only so reranking_score does not land on the caller's + # documents: meta_data and embedding stay shared with the originals. + document = replace(documents[index]) + # The MMR score at the moment this document was picked. Unlike a relevance + # score it is not monotonic across the list, because the candidate pool + # shrinks as redundancy grows: the returned order is authoritative. + document.reranking_score = score + results.append(document) + return results + + def _prepare(self, documents: List[Document], requested: Optional[int] = None) -> Optional[int]: + """Validate inputs and return the number of documents to select.""" + if not documents: + return None + + # A zero vector is indistinguishable from a broken ingest and would otherwise be + # selected as maximally different from everything. + missing = [ + doc.id for doc in documents if doc.embedding is None or len(doc.embedding) == 0 or not any(doc.embedding) + ] + if missing: + # Silently returning the input order would look like MMR ran and found + # nothing to diversify. + raise ValueError( + "MMRReranker requires embeddings on search results, but the vector db returned " + f"{len(missing)} document(s) without one. Some vector dbs (Milvus, MongoDB, " + "Redis, Valkey) do not return embeddings on search." + ) + + # Selecting the whole pool and discarding the tail is wasted work, so stop at + # the count the caller will keep. + candidates = [value for value in (self.top_n, requested) if value is not None] + limit = min(candidates) if candidates else len(documents) + return min(limit, len(documents)) + + def _resolve_embedder(self, documents: List[Document]) -> Any: + """The embedder travels on the search results, so the query uses the indexing model.""" + for document in documents: + if document.embedder is not None: + return document.embedder + raise ValueError( + "MMRReranker needs an embedder to embed the query, but the vector db did not " + "attach one to its search results. Vector dbs that embed queries themselves " + "(such as Upstash hosted embeddings) do not expose one, so MMR cannot run there." + ) + + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + limit = self._prepare(documents, limit) + if limit is None: + return documents + + embedder = self._resolve_embedder(documents) + return self._select(embedder.get_embedding(query), documents, limit) + + async def arerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + selection = self._prepare(documents, limit) + if selection is None: + return documents + + embedder = self._resolve_embedder(documents) + query_embedding = await embedder.async_get_embedding(query) + # Selection is pure-Python and grows with the pool, so keep it off the loop. + return await asyncio.to_thread(self._select, query_embedding, documents, selection) diff --git a/libs/agno/agno/models/aimlapi/__init__.py b/libs/agno/agno/models/aimlapi/__init__.py index 013eafe845b..2234e3dd2d1 100644 --- a/libs/agno/agno/models/aimlapi/__init__.py +++ b/libs/agno/agno/models/aimlapi/__init__.py @@ -1,7 +1,20 @@ -from agno.models.aimlapi.aimlapi import AIMLAPI +from typing import TYPE_CHECKING + from agno.models.aimlapi.constants import AIMLAPI_HEADERS +if TYPE_CHECKING: + from agno.models.aimlapi.aimlapi import AIMLAPI + __all__ = [ "AIMLAPI", "AIMLAPI_HEADERS", ] + + +def __getattr__(name: str): + """Lazy import of the chat model so the attribution constants can be read without `openai` installed.""" + if name == "AIMLAPI": + from agno.models.aimlapi.aimlapi import AIMLAPI + + return AIMLAPI + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/libs/agno/agno/models/anthropic/claude.py b/libs/agno/agno/models/anthropic/claude.py index 033068fd0e1..8e545c6041a 100644 --- a/libs/agno/agno/models/anthropic/claude.py +++ b/libs/agno/agno/models/anthropic/claude.py @@ -1124,14 +1124,14 @@ def _parse_provider_response_delta( response_format: Optional[Union[Dict, Type[BaseModel]]] = None, ) -> ModelResponse: """ - Parse the Claude streaming response into ModelProviderResponse objects. + Parse the Claude streaming response into ModelResponse objects. Args: response: Raw response chunk from Anthropic response_format: Optional response format for structured output parsing Returns: - ModelResponse: Iterator of parsed response data + ModelResponse: Parsed response data """ model_response = ModelResponse() diff --git a/libs/agno/agno/models/base.py b/libs/agno/agno/models/base.py index 77e7f8d9e9d..5898a500093 100644 --- a/libs/agno/agno/models/base.py +++ b/libs/agno/agno/models/base.py @@ -2347,7 +2347,7 @@ def run_function_call( if tool_result.files: function_execution_result.files = tool_result.files else: - function_call_output = str(function_execution_result.result) if function_execution_result.result else "" + function_call_output = str(function_execution_result.result) if function_call.function.show_result and function_call_output is not None: yield ModelResponse(content=function_call_output) diff --git a/libs/agno/agno/models/cohere/chat.py b/libs/agno/agno/models/cohere/chat.py index ab97b7139e3..685917927da 100644 --- a/libs/agno/agno/models/cohere/chat.py +++ b/libs/agno/agno/models/cohere/chat.py @@ -364,9 +364,10 @@ def _parse_provider_response_delta( Args: response: Raw response chunk from the model provider + tool_use: The current tool being built across chunks Returns: - ModelResponse: Parsed response delta + Tuple[ModelResponse, Dict[str, Any]]: The parsed model response delta and updated tool_use """ model_response = ModelResponse() diff --git a/libs/agno/agno/models/groq/groq.py b/libs/agno/agno/models/groq/groq.py index 5ce2d16d7cd..ebd02a78adc 100644 --- a/libs/agno/agno/models/groq/groq.py +++ b/libs/agno/agno/models/groq/groq.py @@ -527,7 +527,7 @@ def _parse_provider_response_delta(self, response: ChatCompletionChunk) -> Model response: Raw response chunk from Groq Returns: - ModelResponse: Iterator of parsed response data + ModelResponse: Parsed response data """ model_response = ModelResponse() diff --git a/libs/agno/agno/models/meta/llama.py b/libs/agno/agno/models/meta/llama.py index 52fc04b77d4..9b6a4c48331 100644 --- a/libs/agno/agno/models/meta/llama.py +++ b/libs/agno/agno/models/meta/llama.py @@ -440,7 +440,7 @@ def _parse_provider_response_delta( Parse the Llama streaming response into a ModelResponse. Args: - response_delta: Raw response chunk from the Llama API + response: Raw response chunk from the Llama API Returns: ModelResponse: Parsed response data diff --git a/libs/agno/agno/models/ollama/chat.py b/libs/agno/agno/models/ollama/chat.py index a0f15ac9520..4920e57030d 100644 --- a/libs/agno/agno/models/ollama/chat.py +++ b/libs/agno/agno/models/ollama/chat.py @@ -404,7 +404,7 @@ def _parse_provider_response_delta(self, response: ChatResponse) -> ModelRespons response (ChatResponse): The response from the provider. Returns: - Iterator[ModelResponse]: An iterator of the model response. + ModelResponse: The parsed response. """ model_response = ModelResponse() diff --git a/libs/agno/agno/models/openai/responses.py b/libs/agno/agno/models/openai/responses.py index 047767cff98..90d5a66446c 100644 --- a/libs/agno/agno/models/openai/responses.py +++ b/libs/agno/agno/models/openai/responses.py @@ -1288,10 +1288,12 @@ def _parse_provider_response_delta( Parse the streaming response from the model provider into a ModelResponse object. Args: - response: Raw response chunk from the model provider + stream_event: Raw streaming event from the model provider + assistant_message: The assistant message to populate + tool_use: The current tool being built across chunks Returns: - ModelResponse: Parsed response delta + Tuple[ModelResponse, Dict[str, Any]]: The parsed model response delta and updated tool_use """ model_response = ModelResponse() @@ -1393,7 +1395,7 @@ def _get_metrics(self, response_usage: ResponseUsage) -> MessageMetrics: Parse the given OpenAI-specific usage into an Agno MessageMetrics object. Args: - response: The response from the provider. + response_usage: Usage data from OpenAI Returns: MessageMetrics: Parsed metrics data diff --git a/libs/agno/agno/models/perplexity/perplexity.py b/libs/agno/agno/models/perplexity/perplexity.py index 49c05ca4ea7..0c6f48c6d28 100644 --- a/libs/agno/agno/models/perplexity/perplexity.py +++ b/libs/agno/agno/models/perplexity/perplexity.py @@ -34,7 +34,7 @@ class Perplexity(OpenAILike): name (str): The model name. Defaults to "Perplexity". provider (str): The provider name. Defaults to "Perplexity". api_key (Optional[str]): The API key. - base_url (str): The base URL. Defaults to "https://api.perplexity.ai/chat/completions". + base_url (str): The base URL. Defaults to "https://api.perplexity.ai/". max_tokens (int): The maximum number of tokens. Defaults to 1024. """ diff --git a/libs/agno/agno/models/utils.py b/libs/agno/agno/models/utils.py index c79c6d2bdf4..a0c837e7f3a 100644 --- a/libs/agno/agno/models/utils.py +++ b/libs/agno/agno/models/utils.py @@ -83,6 +83,7 @@ "xai": ("agno.models.xai", "xAI", "xAI", "xai"), "xai-responses": ("agno.models.xai", "xAIResponses", "xAIResponses", "xai"), "xiaomi": ("agno.models.xiaomi", "MiMo", "MiMo", "xiaomi mimo"), + "yapi": ("agno.models.yapi", "YAPI", "YAPI", "yapi"), } # key -> (module, class_name): the construction registry consumed by `_get_model_class`, the diff --git a/libs/agno/agno/models/vercel/v0.py b/libs/agno/agno/models/vercel/v0.py index e16e973f836..2447f09f855 100644 --- a/libs/agno/agno/models/vercel/v0.py +++ b/libs/agno/agno/models/vercel/v0.py @@ -16,7 +16,7 @@ class V0(OpenAILike): name (str): The name of the API. Defaults to "v0". provider (str): The provider of the API. Defaults to "v0". api_key (Optional[str]): The API key for the v0 API. - base_url (Optional[str]): The base URL for the v0 API. Defaults to "https://v0.dev/chat/settings/keys". + base_url (Optional[str]): The base URL for the v0 API. Defaults to "https://api.v0.dev/v1/". """ id: str = "v0-1.0-md" diff --git a/libs/agno/agno/models/yapi/__init__.py b/libs/agno/agno/models/yapi/__init__.py new file mode 100644 index 00000000000..97c33d2e70b --- /dev/null +++ b/libs/agno/agno/models/yapi/__init__.py @@ -0,0 +1,5 @@ +from agno.models.yapi.yapi import YAPI + +__all__ = [ + "YAPI", +] diff --git a/libs/agno/agno/models/yapi/yapi.py b/libs/agno/agno/models/yapi/yapi.py new file mode 100644 index 00000000000..aa57a273e91 --- /dev/null +++ b/libs/agno/agno/models/yapi/yapi.py @@ -0,0 +1,44 @@ +from dataclasses import dataclass, field +from os import getenv +from typing import Any, Dict, Optional + +from agno.exceptions import ModelAuthenticationError +from agno.models.openai.like import OpenAILike + + +@dataclass +class YAPI(OpenAILike): + """ + A class for interacting with Y-API, an OpenAI-compatible gateway that serves + models from several vendors behind a single endpoint. + + Attributes: + id (str): The id of the Y-API model to use. Default is "deepseek/deepseek-v4-flash". + name (str): The name of this chat model instance. Default is "YAPI". + provider (str): The provider of the model. Default is "YAPI". + api_key (str): The api key to authorize request to Y-API. + base_url (str): The base url to which the requests are sent. + Defaults to "https://api.y-api.bestvirtualgoods.com/v1". + """ + + id: str = "deepseek/deepseek-v4-flash" + name: str = "YAPI" + provider: str = "YAPI" + api_key: Optional[str] = field(default_factory=lambda: getenv("YAPI_API_KEY")) + base_url: str = "https://api.y-api.bestvirtualgoods.com/v1" + + def _get_client_params(self) -> Dict[str, Any]: + """ + Returns client parameters for API requests, checking for YAPI_API_KEY. + + Returns: + Dict[str, Any]: A dictionary of client parameters for API requests. + """ + if not self.api_key: + self.api_key = getenv("YAPI_API_KEY") + if not self.api_key: + raise ModelAuthenticationError( + message="YAPI_API_KEY not set. Please set the YAPI_API_KEY environment variable.", + model_name=self.name, + ) + return super()._get_client_params() diff --git a/libs/agno/agno/os/interfaces/agui/resume.py b/libs/agno/agno/os/interfaces/agui/resume.py index 80e4c1309a2..8f0edcbc9af 100644 --- a/libs/agno/agno/os/interfaces/agui/resume.py +++ b/libs/agno/agno/os/interfaces/agui/resume.py @@ -9,6 +9,7 @@ from agno.session.agent import AgentSession from agno.session.team import TeamSession from agno.team.team import Team +from agno.utils.log import log_warning from agno.utils.string import parse_response_dict_str @@ -39,6 +40,17 @@ def _resolve_external_execution(requirement: RunRequirement, content: str, error requirement.set_external_execution_result(error or content) +def _tool_message_text(tool_message: AGUIToolMessage) -> str: + # ag-ui-protocol 1.0 lets a tool result be a list of content parts; only its text can answer a pause. + content = tool_message.content + if isinstance(content, str): + return content + dropped = sorted({part.type for part in content if part.type != "text"}) + if dropped: + log_warning(f"Tool result {tool_message.tool_call_id}: ignoring {', '.join(dropped)} parts, using its text") + return "\n".join(part.text for part in content if part.type == "text") + + def resolve_requirements_from_tool_messages( requirements: List[RunRequirement], tool_messages: List[AGUIToolMessage], @@ -56,14 +68,15 @@ def resolve_requirements_from_tool_messages( tool_message = tool_message_by_call_id.get(tool_exec.tool_call_id) if tool_message is None: continue + content = _tool_message_text(tool_message) # External execution: raw content, no JSON parsing if requirement.pause_type == "external_execution": - _resolve_external_execution(requirement, tool_message.content, tool_message.error) + _resolve_external_execution(requirement, content, tool_message.error) continue # Structured pause types: parse JSON payload - parsed = parse_response_dict_str(tool_message.content) + parsed = parse_response_dict_str(content) payload: Dict[str, Any] = parsed if isinstance(parsed, dict) else {} if requirement.pause_type == "confirmation": diff --git a/libs/agno/agno/os/routers/knowledge/schemas.py b/libs/agno/agno/os/routers/knowledge/schemas.py index e2c270c5b9f..3436cdefbc1 100644 --- a/libs/agno/agno/os/routers/knowledge/schemas.py +++ b/libs/agno/agno/os/routers/knowledge/schemas.py @@ -158,7 +158,9 @@ class VectorSearchResult(BaseModel): name: Optional[str] = Field(None, description="Name of the document") meta_data: Optional[Dict[str, Any]] = Field(None, description="Metadata associated with the document") usage: Optional[Dict[str, Any]] = Field(None, description="Usage statistics (e.g., token counts)") - reranking_score: Optional[float] = Field(None, description="Reranking score for relevance", ge=0.0, le=1.0) + # Not all rerankers score in [0, 1]: MMR subtracts a redundancy term and goes + # negative, and cross-encoder rerankers write raw logits. + reranking_score: Optional[float] = Field(None, description="Reranking score for relevance", ge=-1.0, le=1.0) content_id: Optional[str] = Field(None, description="ID of the source content") content_origin: Optional[str] = Field(None, description="Origin URL or source of the content") size: Optional[int] = Field(None, description="Size of the content in bytes", ge=0) diff --git a/libs/agno/agno/os/routers/workflows/router.py b/libs/agno/agno/os/routers/workflows/router.py index 3c620b3803e..72d1d438164 100644 --- a/libs/agno/agno/os/routers/workflows/router.py +++ b/libs/agno/agno/os/routers/workflows/router.py @@ -296,13 +296,30 @@ async def handle_workflow_via_websocket( ) return - # Generate session_id if not provided - # Use workflow's default session_id if not provided in message + # A run must not enter a session owned by someone else: the runs table + # has no ownership predicate, so an unguarded write is replayed into + # the owner's history as their own turn. Same guard and same effective + # identity as the HTTP route: the caller's resolved user_id, else the + # workflow's own default, which is what will stamp the session row. + effective_user_id = user_id or getattr(workflow, "user_id", None) + try: + await assert_session_writable( + getattr(workflow, "db", None) or os.db, + session_id, + effective_user_id, + session_type=SessionType.WORKFLOW, + is_admin=bool(ws_auth and ws_auth.is_admin), + ) + except HTTPException as e: + await websocket.send_text(json.dumps({"event": "error", "error": str(e.detail)})) + return + + # A submission that names no session gets a fresh one, as over HTTP. + # The workflow's own session_id is not a default for clients: it + # would pool every client that omits the field into one session, and + # under per-session queueing they would all line up behind each other. if not session_id: - if workflow.session_id: - session_id = workflow.session_id - else: - session_id = str(uuid4()) + session_id = str(uuid4()) # Durable WS submission: the queue row is the acceptance, execution # happens on whichever worker claims it, and this socket becomes a @@ -317,6 +334,10 @@ async def handle_workflow_via_websocket( queue_worker is not None and not is_factory and getattr(workflow, "db", None) is not None + # The worker resolves the registry instance, so a ticket cannot + # carry a version pin: a pinned submission takes the in-process + # path below, where the pin is stamped on the run (as over HTTP) + and version is None and payload_is_queueable(queued_ws_payload) and any( getattr(candidate, "id", None) == workflow_id and not isinstance(candidate, WorkflowFactory) @@ -383,8 +404,8 @@ async def handle_workflow_via_websocket( return if queue_worker is not None: log_warning( - "WS workflow submission bypasses the durable queue (factory/off-registry/no-db " - "workflows are not queueable): bounded and observable, but NOT durable." + "WS workflow submission bypasses the durable queue (factory/off-registry/no-db/" + "version-pinned workflows are not queueable): bounded and observable, but NOT durable." ) # Version-stable preview: an explicitly pinned version is recorded on diff --git a/libs/agno/agno/reasoning/openai.py b/libs/agno/agno/reasoning/openai.py index fee8c062c44..b91461d08c3 100644 --- a/libs/agno/agno/reasoning/openai.py +++ b/libs/agno/agno/reasoning/openai.py @@ -29,6 +29,10 @@ def is_openai_reasoning_model(reasoning_model: Model) -> bool: isinstance(reasoning_model, OpenAILike) and ( "deepseek-r1" in reasoning_model.id.lower() + or "deepseek-reasoner" in reasoning_model.id.lower() + or "deepseek-v3.1" in reasoning_model.id.lower() + or "deepseek-v3.2" in reasoning_model.id.lower() + or "deepseek-v4" in reasoning_model.id.lower() or "minimax-m2" in reasoning_model.id.lower() or "minimax-m3" in reasoning_model.id.lower() ) diff --git a/libs/agno/agno/tools/adanos.py b/libs/agno/agno/tools/adanos.py index dd665250da3..eb8a630d6bd 100644 --- a/libs/agno/agno/tools/adanos.py +++ b/libs/agno/agno/tools/adanos.py @@ -143,7 +143,14 @@ async def aget_stock_sentiment( start_date: Optional[str] = None, end_date: Optional[str] = None, ) -> Dict[str, Any]: - """Asynchronously get sentiment for a stock from one Adanos data source.""" + """Get sentiment for a stock from one Adanos data source. + + Args: + ticker: Stock ticker, for example ``AAPL`` or ``TSLA``. + source: Sentiment source: reddit, x, news, or polymarket. + start_date: Inclusive UTC start date in YYYY-MM-DD format. + end_date: Inclusive UTC end date in YYYY-MM-DD format. + """ path = self._STOCK_PATHS.get(source) if path is None: return {"error": "source must be one of: reddit, x, news, polymarket"} @@ -166,7 +173,13 @@ def get_crypto_sentiment( async def aget_crypto_sentiment( self, symbol: str, start_date: Optional[str] = None, end_date: Optional[str] = None ) -> Dict[str, Any]: - """Asynchronously get Reddit sentiment for a cryptocurrency.""" + """Get Reddit sentiment for a cryptocurrency. + + Args: + symbol: Cryptocurrency symbol, for example ``BTC`` or ``ETH``. + start_date: Inclusive UTC start date in YYYY-MM-DD format. + end_date: Inclusive UTC end date in YYYY-MM-DD format. + """ normalized_symbol = quote(symbol.strip().upper(), safe=".-") return await self._arequest( f"{self._CRYPTO_PATH}/token/{normalized_symbol}", self._params(start_date, end_date) @@ -203,7 +216,15 @@ async def aget_trending( start_date: Optional[str] = None, end_date: Optional[str] = None, ) -> Dict[str, Any]: - """Asynchronously get trending stocks or cryptocurrencies ranked by buzz score with sentiment data.""" + """Get trending stocks or cryptocurrencies ranked by buzz score with sentiment data. + + Args: + asset_type: Asset universe: stocks or crypto. + source: For stocks, reddit, x, news, or polymarket. Crypto uses reddit. + limit: Maximum number of results, from 1 to 100. + start_date: Inclusive UTC start date in YYYY-MM-DD format. + end_date: Inclusive UTC end date in YYYY-MM-DD format. + """ path = self._asset_path(asset_type, source) if isinstance(path, dict): return path @@ -237,7 +258,14 @@ async def aget_market_sentiment( start_date: Optional[str] = None, end_date: Optional[str] = None, ) -> Dict[str, Any]: - """Asynchronously get aggregate market sentiment for stocks or cryptocurrencies.""" + """Get aggregate market sentiment for stocks or cryptocurrencies. + + Args: + asset_type: Asset universe: stocks or crypto. + source: For stocks, reddit, x, news, or polymarket. Crypto uses reddit. + start_date: Inclusive UTC start date in YYYY-MM-DD format. + end_date: Inclusive UTC end date in YYYY-MM-DD format. + """ path = self._asset_path(asset_type, source) if isinstance(path, dict): return path diff --git a/libs/agno/agno/tools/calcom.py b/libs/agno/agno/tools/calcom.py index 217d6d757b2..e9e89be7ce1 100644 --- a/libs/agno/agno/tools/calcom.py +++ b/libs/agno/agno/tools/calcom.py @@ -71,7 +71,6 @@ def _convert_to_user_timezone(self, utc_time: str) -> str: Args: utc_time: UTC time string - user_timezone: User's timezone (e.g., 'Asia/Kolkata') Returns: str: Formatted time in user's timezone @@ -106,8 +105,6 @@ def get_available_slots( Args: start_date: Start date in YYYY-MM-DD format end_date: End date in YYYY-MM-DD format - user_timezone: User's timezone - event_type_id: Optional specific event type ID Returns: str: Available slots or error message @@ -214,7 +211,6 @@ def reschedule_booking( booking_uid: Booking UID to reschedule new_start_time: New start time in YYYY-MM-DDTHH:MM:SSZ format reason: Reason for rescheduling - user_timezone: User's timezone Returns: str: Rescheduling confirmation or error message diff --git a/libs/agno/agno/tools/coding.py b/libs/agno/agno/tools/coding.py index 6d11c4f5edf..fe016838dce 100644 --- a/libs/agno/agno/tools/coding.py +++ b/libs/agno/agno/tools/coding.py @@ -14,19 +14,37 @@ @functools.lru_cache(maxsize=None) def _warn_coding_tools() -> None: - logger.warning("CodingTools can run arbitrary shell commands, please provide human supervision.") + logger.warning( + "CodingTools run_shell executes arbitrary shell commands. Provide human supervision " + "and never expose it to untrusted input; restrict_to_base_dir is not a security sandbox." + ) class CodingTools(Toolkit): """A minimal, powerful toolkit for coding agents. - Provides four core tools (read, edit, write, shell) and three optional - exploration tools (grep, find, ls). With these primitives, an agent can + Provides three core tools (read, edit, write) plus opt-in shell (run_shell) + and exploration tools (grep, find, ls). With these primitives, an agent can perform any file operation, run tests, use git, install packages, search codebases, and more. Inspired by the Pi coding agent's philosophy: a small number of composable tools is more powerful than many specialized ones. + + Security: + run_shell executes commands through the system shell and is disabled by + default. Enable it only for agents you supervise, and never expose it to + untrusted or third-party input. + + ``restrict_to_base_dir`` reduces accidental damage: it confines file + tools to base_dir and, for run_shell, runs commands without a shell + (so chaining, redirection, substitution, and globbing are inert) behind + a command allowlist. It is NOT a security sandbox. Any allowlisted + interpreter (python, pip, git, ...) can read arbitrary files, dump the + process environment, or reach the network, so a determined caller can + escape the restriction. To run shell against untrusted input, execute it + in a real sandbox (separate process or container with a scrubbed + environment, no network, and a read-only mount) instead. """ DEFAULT_ALLOWED_COMMANDS: List[str] = [ @@ -129,7 +147,7 @@ def __init__( enable_read_file: bool = True, enable_edit_file: bool = True, enable_write_file: bool = True, - enable_run_shell: bool = True, + enable_run_shell: bool = False, enable_grep: bool = False, enable_find: bool = False, enable_ls: bool = False, @@ -143,14 +161,20 @@ def __init__( Args: base_dir: Root directory for file operations. Defaults to cwd. - restrict_to_base_dir: If True, file and shell operations cannot escape base_dir. + restrict_to_base_dir: If True, confine file tools to base_dir and run + run_shell commands without a shell behind a command allowlist (so + chaining, redirection, substitution, and globbing are inert). This + limits accidental damage but is not a security sandbox: an allowlisted + interpreter can still escape it. Do not rely on it for untrusted input. + If False, run_shell runs the raw string through the system shell. max_lines: Maximum lines to return before truncating (default 2000). max_bytes: Maximum bytes to return before truncating (default 50KB). shell_timeout: Timeout in seconds for shell commands (default 120). enable_read_file: Enable the read_file tool. enable_edit_file: Enable the edit_file tool. enable_write_file: Enable the write_file tool. - enable_run_shell: Enable the run_shell tool. + enable_run_shell: Enable the run_shell tool. Disabled by default because it + executes arbitrary shell commands; enable it only under human supervision. enable_grep: Enable the grep tool (disabled by default). enable_find: Enable the find tool (disabled by default). enable_ls: Enable the ls tool (disabled by default). @@ -247,39 +271,104 @@ def _cleanup_temp_files(self) -> None: pass self._temp_files.clear() - # Shell operators that enable command chaining or substitution - _DANGEROUS_PATTERNS: List[str] = ["&&", "||", ";", "|", "$(", "`", ">", ">>", "<"] + # Control operators that chain, background, or redirect commands. In restricted + # mode commands run without a shell (shell=False), so these never take effect; + # shlex leaves an unquoted operator as its own token, which we reject with a + # clear "unsupported" message while quoted uses (e.g. -m "A & B") pass untouched. + _UNSUPPORTED_OPERATOR_TOKENS: set = {"&&", "||", ";", "|", "&", "<", ">", ">>"} + + # Interpreters that can execute arbitrary inline code, bypassing the allowlist + # and path checks. Matched by basename prefix (python, python3, python3.12, ...). + _CODE_EXEC_INTERPRETER_PREFIXES: tuple = ("python",) + + # CPython short options that execute arbitrary inline code (-c cmd, -m module). + _CODE_EXEC_SHORT_OPTS: set = {"c", "m"} + + # CPython short options that consume the rest of the token as their argument, so + # a following 'c'/'m' is a value, not the code-exec flag (e.g. -W c, -X c). + _ARG_TAKING_SHORT_OPTS: set = {"W", "X", "Q"} + + def _has_interpreter_code_exec(self, args: List[str]) -> bool: + """Detect inline code execution in a Python interpreter's arguments. + + Handles attached and clustered short options the way CPython does, e.g. + ``-c``, ``-c'code'``, ``-mmod``, ``-Ic 'code'``. Option parsing stops at + the first non-option argument (the script path), and short options that + take a value (-W, -X, -Q) consume the remainder of their token, so a 'c' + or 'm' appearing as such a value is not treated as code execution. + """ + for token in args: + if token == "-": # program read from stdin + return True + if not token.startswith("-"): + # First positional is the script path; CPython stops parsing options here. + break + if token.startswith("--"): + # No CPython long option executes inline code. + continue + for ch in token[1:]: + if ch in self._CODE_EXEC_SHORT_OPTS: + return True + if ch in self._ARG_TAKING_SHORT_OPTS: + # Remainder of this token is the option's argument, not more flags. + break + return False def _check_command(self, command: str) -> Optional[str]: - """Check if a shell command is safe to execute. + """Validate a command for restricted mode, returning an error message or None. - When restrict_to_base_dir is True, this method: - 1. Blocks shell metacharacters that enable chaining/substitution. + In restricted mode the command is executed without a shell (see run_shell), + so this validates the shlex-tokenized command that will actually run: + 1. Rejects control operators (|, &&, ;, redirects) as unsupported, since a + shell-less run would treat them as literal arguments, not chaining. 2. Validates the command name against the allowed_commands list (if set). - 3. Checks that path-like tokens don't escape the base directory. + 3. Blocks inline code-execution flags on interpreters (e.g. python3 -c). + 4. Checks that path-like tokens don't escape the base directory. - Returns an error message if a violation is found, None if safe. + Allowlist and operator rejection are harm reduction, not a security sandbox; + running without a shell is what actually neutralizes chaining/substitution. """ if not self.restrict_to_base_dir: return None - # Block shell operators that enable chaining/substitution - for pattern in self._DANGEROUS_PATTERNS: - if pattern in command: - return f"Error: Shell operator '{pattern}' is not allowed in restricted mode." - try: tokens = shlex.split(command) except ValueError: return "Error: Could not parse shell command." + # Reject control operators to give a clear error instead of a confusing literal + # run (shell=False already makes them inert). shlex leaves an unquoted operator + # as its own token while keeping it inside a larger quoted argument, so + # `git commit -m "A & B"` passes. Known limitation: an argument that is exactly + # an operator (e.g. `echo '&'`) also becomes a bare token and is rejected; that + # is harmless over-rejection, and detecting it would require full quote tracking. + for token in tokens: + if token in self._UNSUPPORTED_OPERATOR_TOKENS: + return ( + f"Error: Shell operator '{token}' is not supported in restricted mode. " + "Run separate commands, or set restrict_to_base_dir=False for a full " + "shell (trusted, supervised use only)." + ) + # Validate command against allowlist + cmd_base = Path(tokens[0]).name if tokens else "" # Handle /usr/bin/python -> python if self.allowed_commands is not None and tokens: - cmd = tokens[0] - cmd_base = Path(cmd).name # Handle /usr/bin/python -> python if cmd_base not in self.allowed_commands: return f"Error: Command '{cmd_base}' is not in the allowed commands list." + # Block inline code execution via an interpreter, which would otherwise run + # arbitrary code past the allowlist and path checks (e.g. python3 -c "..."). + # This is harm reduction, not a boundary: an interpreter can still escape by + # running a script file. Do not expose run_shell to untrusted input. + if cmd_base.startswith(self._CODE_EXEC_INTERPRETER_PREFIXES): + if self._has_interpreter_code_exec(tokens[1:]): + return ( + "Error: Inline code execution (-c/-m or reading from stdin) is not " + "allowed in restricted mode. Run a script file instead. Setting " + "restrict_to_base_dir=False lifts all checks and should only be used " + "for trusted, supervised execution." + ) + for i, token in enumerate(tokens): # Skip the command itself (already validated by allowlist above) if i == 0: @@ -497,14 +586,17 @@ def write_file(self, file_path: str, contents: str) -> str: return f"Error writing file: {e}" def run_shell(self, command: str, timeout: Optional[int] = None) -> str: - """Execute a shell command and return its output. + """Execute a command and return its output. - Runs the command as a string via the system shell. Output (stdout + stderr) - is truncated if it exceeds the configured limits. When output is truncated, - the full output is saved to a temporary file and its path is included in - the response. + In restricted mode the command is tokenized and run WITHOUT a shell, so + chaining, redirection, command substitution, and globbing have no effect; + this is what makes the allowlist meaningful. When restrict_to_base_dir is + False the raw string is run through the system shell instead (full power, + for trusted and supervised use only). Output (stdout + stderr) is truncated + if it exceeds the configured limits, with the full output saved to a temp + file whose path is included in the response. - :param command: The shell command to execute as a single string. + :param command: The command to execute as a single string. :param timeout: Timeout in seconds. Defaults to the toolkit's shell_timeout. :return: Command output (stdout and stderr combined), or an error message. """ @@ -512,16 +604,30 @@ def run_shell(self, command: str, timeout: Optional[int] = None) -> str: _warn_coding_tools() log_info(f"Running shell command: {command}") - # Check for path escapes in command - path_error = self._check_command(command) - if path_error: - return path_error + # Validate against the restricted-mode policy (allowlist, operators, paths). + command_error = self._check_command(command) + if command_error: + return command_error effective_timeout = timeout if timeout is not None else self.shell_timeout + # Restricted mode runs without a shell so operators cannot chain or + # substitute; unrestricted mode keeps full shell semantics by request. + if self.restrict_to_base_dir: + try: + args: Union[str, List[str]] = shlex.split(command) + except ValueError: + return "Error: Could not parse shell command." + if not args: + return "Error: Empty command." + use_shell = False + else: + args = command + use_shell = True + result = subprocess.run( - command, - shell=True, + args, + shell=use_shell, capture_output=True, text=True, timeout=effective_timeout, @@ -556,6 +662,9 @@ def run_shell(self, command: str, timeout: Optional[int] = None) -> str: except subprocess.TimeoutExpired: effective_timeout = timeout if timeout is not None else self.shell_timeout return f"Error: Command timed out after {effective_timeout} seconds" + except FileNotFoundError: + # Raised in restricted mode (shell=False) when the executable is missing. + return f"Error: Command not found: {command}" except Exception as e: log_error(f"Error running shell command: {str(e)}") return f"Error running shell command: {e}" diff --git a/libs/agno/agno/tools/github.py b/libs/agno/agno/tools/github.py index 044122e65e5..ec5712d4535 100644 --- a/libs/agno/agno/tools/github.py +++ b/libs/agno/agno/tools/github.py @@ -710,7 +710,7 @@ def get_pull_requests( state (str, optional): State of the PRs to retrieve. Can be 'open', 'closed', or 'all'. Defaults to 'open'. sort (str, optional): What to sort results by. Can be 'created', 'updated', 'popularity', 'long-running'. Defaults to 'created'. direction (str, optional): The direction of the sort. Can be 'asc' or 'desc'. Defaults to 'desc'. - limit (int, optional): The maximum number of pull requests to return. Defaults to 20. + limit (int, optional): The maximum number of pull requests to return. Defaults to 50. Returns: A JSON-formatted string containing a list of pull requests. diff --git a/libs/agno/agno/tools/google/bigquery.py b/libs/agno/agno/tools/google/bigquery.py index bb28121eb6b..0176a7beb97 100644 --- a/libs/agno/agno/tools/google/bigquery.py +++ b/libs/agno/agno/tools/google/bigquery.py @@ -83,7 +83,7 @@ def list_tables(self) -> str: def describe_table(self, table_id: str) -> str: """Use this function to describe a table. Args: - table_name (str): The name of the table to get the schema for. + table_id (str): The ID of the table to get the schema for. Returns: str: schema of a table """ diff --git a/libs/agno/agno/tools/google/gmail.py b/libs/agno/agno/tools/google/gmail.py index 55dac947280..5a2c59d54a0 100644 --- a/libs/agno/agno/tools/google/gmail.py +++ b/libs/agno/agno/tools/google/gmail.py @@ -1501,10 +1501,10 @@ def search_threads(self, query: str = "", count: int = 10, page_token: Optional[ Args: query: Gmail search query string. Supports all Gmail operators like from:, to:, subject:, is:unread, etc. count: Maximum number of threads to return (default 10, max 500). - next_page_token: Token for pagination. + page_token: Token from a previous response to fetch the next page. Returns: - JSON string with list of matching threads and next_page_token if more results exist. + JSON string with list of matching threads and nextPageToken if more results exist. """ try: service = self.service @@ -1623,10 +1623,10 @@ def list_drafts(self, count: int = 10, page_token: Optional[str] = None) -> str: Args: count: Maximum number of drafts to return (default 10, max 500). - next_page_token: Token for pagination. + page_token: Token from a previous response to fetch the next page. Returns: - JSON string with list of draft IDs and next_page_token if more results exist. + JSON string with list of draft IDs and nextPageToken if more results exist. """ try: service = self.service diff --git a/libs/agno/agno/tools/mcp/mcp.py b/libs/agno/agno/tools/mcp/mcp.py index 5b41b9bb820..055824384f2 100644 --- a/libs/agno/agno/tools/mcp/mcp.py +++ b/libs/agno/agno/tools/mcp/mcp.py @@ -20,10 +20,14 @@ try: from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import get_default_environment + from mcp.types import ListToolsResult, PaginatedRequestParams except ModuleNotFoundError: raise ImportError("`mcp` not installed. Please install using `pip install 'mcp>=2.1.0,<3.0.0'`") +# Match the default bound on automatic discovery through fastmcp's Client. +_MCP_TOOL_PAGINATION_MAX_PAGES = 250 + _FASTMCP_INSTALL_HINT = ( "`fastmcp` not installed. MCPTools builds its connections with it. " "Please install using `pip install 'fastmcp>=4.0.0,<5'`" @@ -889,7 +893,17 @@ async def build_tools(self) -> None: listed = await self.session.list_tools() # fastmcp's Client yields a plain list; a user-supplied ClientSession # yields a ListToolsResult carrying .tools. - available_tools = listed if isinstance(listed, list) else listed.tools + available_tools = list(listed if isinstance(listed, list) else listed.tools) + if isinstance(listed, ListToolsResult): + # Collect all pages before validating filters or changing the registry. + page_count = 1 + while listed.next_cursor is not None: + if page_count >= _MCP_TOOL_PAGINATION_MAX_PAGES: + raise RuntimeError(f"MCP tools/list reached the page limit ({_MCP_TOOL_PAGINATION_MAX_PAGES})") + # Cursors are opaque: empty or repeated values may still advance the listing. + listed = await self.session.list_tools(params=PaginatedRequestParams(cursor=listed.next_cursor)) + available_tools.extend(listed.tools) + page_count += 1 self._check_tools_filters( available_tools=[tool.name for tool in available_tools], diff --git a/libs/agno/agno/tools/minimax.py b/libs/agno/agno/tools/minimax.py index 8423836e754..f81bab44df6 100644 --- a/libs/agno/agno/tools/minimax.py +++ b/libs/agno/agno/tools/minimax.py @@ -80,7 +80,7 @@ def generate_video( Args: prompt: Text description of the video to generate. - resolution: Output resolution. MiniMax H3 currently supports 2K. + resolution: Output resolution. MiniMax H3 supports 768P or 2K. duration: Video duration in seconds, from 4 through 15. ratio: Output aspect ratio, such as 16:9 or 9:16. """ @@ -162,7 +162,14 @@ async def agenerate_video( duration: int = 5, ratio: str = "16:9", ) -> ToolResult: - """Generate a video from a text prompt asynchronously.""" + """Generate a video from a text prompt. + + Args: + prompt: Text description of the video to generate. + resolution: Output resolution. MiniMax H3 supports 768P or 2K. + duration: Video duration in seconds, from 4 through 15. + ratio: Output aspect ratio, such as 16:9 or 9:16. + """ if not self.api_key: return ToolResult(content="Please set the MINIMAX_API_KEY") if not prompt: diff --git a/libs/agno/agno/tools/models/aimlapi.py b/libs/agno/agno/tools/models/aimlapi.py new file mode 100644 index 00000000000..7237fb00b31 --- /dev/null +++ b/libs/agno/agno/tools/models/aimlapi.py @@ -0,0 +1,658 @@ +import asyncio +import mimetypes +import re +import time +from os import getenv +from pathlib import Path +from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, Union +from urllib.parse import urlsplit +from uuid import uuid4 + +import httpx + +from agno.media import Audio, Image, Video +from agno.models.aimlapi.constants import AIMLAPI_HEADERS +from agno.tools import Toolkit +from agno.tools.function import ToolResult +from agno.utils.log import log_debug, log_error, log_warning + +DEFAULT_BASE_URL = "https://api.aimlapi.com" +# The attribution headers mean something only on this host, so a proxy or a +# self-hosted mirror in front of the API is sent none of them. +AIMLAPI_HOST = "api.aimlapi.com" + +SPEECH_FORMATS = ("mp3", "opus", "aac", "flac", "wav", "pcm") + +# Statuses a job reports while it is still running. Anything outside this set is +# terminal: "completed" carries the result, "error"/"failed" carry a message, and +# an unknown status stops the loop instead of spinning until the timeout. +# Measured on the gateway 2026-09-23 across six transcription models: only +# queued, generating, completed and failed appear. "waiting" is kept because the +# Nova-3 docs example still tests for it, and treating it as in-progress can only +# ever mean one more poll. +_IN_PROGRESS = frozenset( + {"queued", "generating", "processing", "pending", "running", "in_progress", "active", "waiting"} +) +# Deepgram-backed jobs report "error"; AssemblyAI-backed ones report "failed". +_FAILED = frozenset({"error", "failed"}) +_TRANSIENT_STATUSES = frozenset({408, 425, 429, 500, 502, 503, 504}) +_MAX_TRANSIENT_RETRIES = 3 + + +class AIMLAPIError(RuntimeError): + """The gateway answered with an error status.""" + + def __init__(self, status_code: int, message: str): + super().__init__(f"AI/ML API returned HTTP {status_code}: {message}") + self.status_code = status_code + + +class AIMLAPITools(Toolkit): + """Tools for the media endpoints of AI/ML API (https://aimlapi.com). + + One key gives an agent image, video, speech and transcription models from + many vendors behind one endpoint. Each capability is a separate tool with + its own model, so an agent can be given only the ones it needs. Every tool + has an async variant, so the long-running video and transcription jobs do + not block the event loop under ``arun``. + + Args: + api_key (str, optional): AI/ML API key. Read from AIMLAPI_API_KEY if not provided. + base_url (str): API root. Default is "https://api.aimlapi.com". A trailing "/v1" + (the form the AIMLAPI chat model uses) is accepted and stripped. + enable_generate_image (bool): Register generate_image. Default is True. + enable_generate_video (bool): Register generate_video. Default is True. + enable_generate_speech (bool): Register generate_speech. Default is True. + enable_transcribe_audio (bool): Register transcribe_audio. Default is True. + all (bool): Register every tool, overriding the individual flags. Default is False. + image_model (str): Image model id. Default is "openai/gpt-image-2". + image_size (str, optional): "WIDTHxHEIGHT" when the model takes one. + image_quality (str, optional): Quality preset when the model takes one. + video_model (str): Video model id. Default is "bytedance/seedance-2-5". + video_duration (int, optional): Clip length in seconds when the model takes one. + video_resolution (str, optional): E.g. "720p" when the model takes one. + video_aspect_ratio (str, optional): E.g. "16:9" when the model takes one. + video_poll_interval (float): Seconds between status checks. Default is 5. + video_timeout (float): Seconds to wait for a video before giving up. Default is 900. + speech_model (str): Text-to-speech model id. Default is "openai/tts-1". + speech_voice (str, optional): Voice name when the model takes one. Default is "alloy". + speech_format (str): Output container: mp3, opus, aac, flac, wav or pcm. Default is "mp3". + speech_speed (float, optional): Playback speed multiplier when the model takes one. + transcription_model (str): Speech-to-text model id. Default is "deepgram/nova-3". + transcription_language (str, optional): Language hint when the model takes one. + transcription_poll_interval (float): Seconds between status checks. Default is 2. + transcription_timeout (float): Seconds to wait for a transcript. Default is 300. + base_dir (Path or str, optional): Directory local audio files for transcription are + read from. Default is the current working directory. + restrict_to_base_dir (bool): Refuse local paths that resolve outside base_dir, so a + prompt cannot make the agent upload an arbitrary file. Default is True. + timeout (int): Seconds allowed for one HTTP call. Default is 120. + """ + + def __init__( + self, + api_key: Optional[str] = None, + base_url: str = DEFAULT_BASE_URL, + enable_generate_image: bool = True, + enable_generate_video: bool = True, + enable_generate_speech: bool = True, + enable_transcribe_audio: bool = True, + all: bool = False, + image_model: str = "openai/gpt-image-2", + image_size: Optional[str] = None, + image_quality: Optional[str] = None, + video_model: str = "bytedance/seedance-2-5", + video_duration: Optional[int] = None, + video_resolution: Optional[str] = None, + video_aspect_ratio: Optional[str] = None, + video_poll_interval: float = 5.0, + video_timeout: float = 900.0, + speech_model: str = "openai/tts-1", + speech_voice: Optional[str] = "alloy", + speech_format: str = "mp3", + speech_speed: Optional[float] = None, + transcription_model: str = "deepgram/nova-3", + transcription_language: Optional[str] = None, + transcription_poll_interval: float = 2.0, + transcription_timeout: float = 300.0, + base_dir: Optional[Union[Path, str]] = None, + restrict_to_base_dir: bool = True, + timeout: int = 120, + **kwargs, + ): + self.api_key = api_key or getenv("AIMLAPI_API_KEY") + if not self.api_key: + raise ValueError("AIMLAPI_API_KEY not set. Please set the AIMLAPI_API_KEY environment variable.") + if speech_format not in SPEECH_FORMATS: + raise ValueError(f"speech_format must be one of {', '.join(SPEECH_FORMATS)}, got {speech_format!r}") + + # The chat model's base URL ends in /v1; this toolkit addresses both + # /v1 and /v2 routes, so it wants the bare root. + self.base_url = re.sub(r"/v\d+/?$", "", base_url.rstrip("/")) + self.image_model = image_model + self.image_size = image_size + self.image_quality = image_quality + self.video_model = video_model + self.video_duration = video_duration + self.video_resolution = video_resolution + self.video_aspect_ratio = video_aspect_ratio + self.video_poll_interval = video_poll_interval + self.video_timeout = video_timeout + self.speech_model = speech_model + self.speech_voice = speech_voice + self.speech_format = speech_format + self.speech_speed = speech_speed + self.transcription_model = transcription_model + self.transcription_language = transcription_language + self.transcription_poll_interval = transcription_poll_interval + self.transcription_timeout = transcription_timeout + self.base_dir = Path(base_dir) if base_dir is not None else Path.cwd() + self.restrict_to_base_dir = restrict_to_base_dir + self.request_timeout = timeout + + tools: List[Any] = [] + async_tools: List[Tuple[Callable[..., Any], str]] = [] + if all or enable_generate_image: + tools.append(self.generate_image) + async_tools.append((self.agenerate_image, "generate_image")) + if all or enable_generate_video: + tools.append(self.generate_video) + async_tools.append((self.agenerate_video, "generate_video")) + if all or enable_generate_speech: + tools.append(self.generate_speech) + async_tools.append((self.agenerate_speech, "generate_speech")) + if all or enable_transcribe_audio: + tools.append(self.transcribe_audio) + async_tools.append((self.atranscribe_audio, "transcribe_audio")) + + super().__init__(name="aimlapi_tools", tools=tools, async_tools=async_tools, timeout=timeout, **kwargs) + + # --- HTTP --------------------------------------------------------------- + + def _headers(self) -> Dict[str, str]: + headers = {"Authorization": f"Bearer {self.api_key}"} + if urlsplit(self.base_url).hostname == AIMLAPI_HOST: + headers.update(AIMLAPI_HEADERS) + return headers + + @staticmethod + def _json(response: httpx.Response) -> Dict[str, Any]: + if response.status_code >= 400: + raise AIMLAPIError(response.status_code, _error_message(response)) + body = response.json() + if not isinstance(body, dict): + raise RuntimeError("AI/ML API returned a non-object JSON body") + return body + + def _post(self, path: str, body: Dict[str, Any]) -> Dict[str, Any]: + response = httpx.post( + f"{self.base_url}{path}", json=body, headers=self._headers(), timeout=self.request_timeout + ) + return self._json(response) + + def _post_file(self, path: str, data: Dict[str, Any], file: Path) -> Dict[str, Any]: + with file.open("rb") as handle: + response = httpx.post( + f"{self.base_url}{path}", + data=data, + files={"audio": (file.name, handle)}, + headers=self._headers(), + timeout=self.request_timeout, + ) + return self._json(response) + + def _get(self, path: str, params: Optional[Dict[str, str]] = None) -> Dict[str, Any]: + response = httpx.get( + f"{self.base_url}{path}", params=params, headers=self._headers(), timeout=self.request_timeout + ) + return self._json(response) + + def _download(self, url: str) -> Tuple[bytes, str]: + """Fetch a generated asset. The link is public, so no key travels with it.""" + parsed = urlsplit(url) + if parsed.scheme != "https": + raise RuntimeError("AI/ML API returned a non-HTTPS asset URL") + response = httpx.get(url, follow_redirects=True, timeout=self.request_timeout) + response.raise_for_status() + return response.content, _asset_mime_type(response.headers.get("content-type"), url) + + async def _apost(self, path: str, body: Dict[str, Any]) -> Dict[str, Any]: + async with httpx.AsyncClient(timeout=self.request_timeout) as client: + response = await client.post(f"{self.base_url}{path}", json=body, headers=self._headers()) + return self._json(response) + + async def _apost_file(self, path: str, data: Dict[str, Any], file: Path) -> Dict[str, Any]: + async with httpx.AsyncClient(timeout=self.request_timeout) as client: + with file.open("rb") as handle: + response = await client.post( + f"{self.base_url}{path}", data=data, files={"audio": (file.name, handle)}, headers=self._headers() + ) + return self._json(response) + + async def _aget(self, path: str, params: Optional[Dict[str, str]] = None) -> Dict[str, Any]: + async with httpx.AsyncClient(timeout=self.request_timeout) as client: + response = await client.get(f"{self.base_url}{path}", params=params, headers=self._headers()) + return self._json(response) + + async def _adownload(self, url: str) -> Tuple[bytes, str]: + parsed = urlsplit(url) + if parsed.scheme != "https": + raise RuntimeError("AI/ML API returned a non-HTTPS asset URL") + async with httpx.AsyncClient(timeout=self.request_timeout, follow_redirects=True) as client: + response = await client.get(url) + response.raise_for_status() + return response.content, _asset_mime_type(response.headers.get("content-type"), url) + + # --- polling ------------------------------------------------------------ + + def _poll(self, fetch: Callable[[], Dict[str, Any]], interval: float, timeout: float, what: str) -> Dict[str, Any]: + """Poll a job until it leaves the in-progress statuses. + + A transient gateway or network error during a poll does not abandon the + job, which keeps running (and billing) on the other side: the poll is + retried a few times before giving up. + """ + deadline = time.monotonic() + timeout + failures = 0 + while True: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"{what} still running after {timeout:.0f}s") + time.sleep(min(interval, remaining)) + try: + job = fetch() + except (AIMLAPIError, httpx.TransportError) as e: + if not _is_transient(e) or failures >= _MAX_TRANSIENT_RETRIES: + raise + failures += 1 + log_warning(f"{what}: poll failed ({e}); retry {failures}/{_MAX_TRANSIENT_RETRIES}") + continue + failures = 0 + if job.get("status") not in _IN_PROGRESS: + return job + + async def _apoll( + self, fetch: Callable[[], Awaitable[Dict[str, Any]]], interval: float, timeout: float, what: str + ) -> Dict[str, Any]: + deadline = time.monotonic() + timeout + failures = 0 + while True: + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError(f"{what} still running after {timeout:.0f}s") + await asyncio.sleep(min(interval, remaining)) + try: + job = await fetch() + except (AIMLAPIError, httpx.TransportError) as e: + if not _is_transient(e) or failures >= _MAX_TRANSIENT_RETRIES: + raise + failures += 1 + log_warning(f"{what}: poll failed ({e}); retry {failures}/{_MAX_TRANSIENT_RETRIES}") + continue + failures = 0 + if job.get("status") not in _IN_PROGRESS: + return job + + # --- request and response shapes ---------------------------------------- + + def _image_body(self, prompt: str) -> Dict[str, Any]: + body: Dict[str, Any] = {"model": self.image_model, "prompt": prompt} + if self.image_size: + body["size"] = self.image_size + if self.image_quality: + body["quality"] = self.image_quality + return body + + def _video_body(self, prompt: str) -> Dict[str, Any]: + body: Dict[str, Any] = {"model": self.video_model, "prompt": prompt} + if self.video_duration is not None: + body["duration"] = self.video_duration + if self.video_resolution: + body["resolution"] = self.video_resolution + if self.video_aspect_ratio: + body["aspect_ratio"] = self.video_aspect_ratio + return body + + def _speech_body(self, text_input: str) -> Dict[str, Any]: + body: Dict[str, Any] = {"model": self.speech_model, "text": text_input, "response_format": self.speech_format} + if self.speech_voice: + body["voice"] = self.speech_voice + if self.speech_speed is not None: + body["speed"] = self.speech_speed + return body + + def _transcription_data(self) -> Dict[str, Any]: + data: Dict[str, Any] = {"model": self.transcription_model} + if self.transcription_language: + data["language"] = self.transcription_language + return data + + def _local_audio(self, audio_path: str) -> Path: + """Resolve a local path inside base_dir; the model must not pick arbitrary files.""" + safe, resolved = self._check_path(audio_path, self.base_dir, self.restrict_to_base_dir) + if not safe: + raise PermissionError(f"{audio_path} is outside the allowed directory {self.base_dir}") + if not resolved.is_file(): + raise FileNotFoundError(f"{audio_path} is not a file") + return resolved + + @staticmethod + def _image_result(prompt: str, assets: List[Tuple[bytes, str]]) -> ToolResult: + if not assets: + log_warning("AI/ML API returned no image data.") + return ToolResult(content="Failed to generate image: No image data received from API.") + images = [ + Image( + id=str(uuid4()), + content=content, + mime_type=mime_type, + format=_subtype(mime_type), + original_prompt=prompt, + ) + for content, mime_type in assets + ] + log_debug(f"Generated {len(images)} image(s)") + return ToolResult(content="Image generated successfully.", images=images) + + @staticmethod + def _video_result(prompt: str, content: bytes, mime_type: str) -> ToolResult: + video = Video( + id=str(uuid4()), content=content, mime_type=mime_type, format=_subtype(mime_type), original_prompt=prompt + ) + log_debug(f"Generated video {video.id} ({len(content)} bytes)") + return ToolResult(content="Video generated successfully.", videos=[video]) + + def _speech_result(self, content: bytes, mime_type: str) -> ToolResult: + audio = Audio(id=str(uuid4()), content=content, mime_type=mime_type, format=self.speech_format) + return ToolResult(content=f"Speech generated successfully with ID: {audio.id}", audios=[audio]) + + @staticmethod + def _job_failure(job: Dict[str, Any], what: str) -> Optional[str]: + """A message when a finished job did not succeed, else None.""" + status = job.get("status") + if status == "completed": + return None + if status in _FAILED: + return f"Failed to {what}: {_error_text(job.get('error')) or 'generation failed'}" + return f"Failed to {what}: job ended with status {status!r}" + + @staticmethod + def _transcript(job: Dict[str, Any]) -> Optional[str]: + """The transcript out of a completed job. Providers differ in where they put it.""" + result = job.get("result") + if not isinstance(result, dict): + return None + for key in ("text", "transcript"): + if isinstance(result.get(key), str): + return result[key] + results = result.get("results") + channels = results.get("channels") if isinstance(results, dict) else None + for channel in channels or []: + if not isinstance(channel, dict): + continue + for alternative in channel.get("alternatives") or []: + if isinstance(alternative, dict) and isinstance(alternative.get("transcript"), str): + return alternative["transcript"] + return None + + # --- tools -------------------------------------------------------------- + + def generate_image(self, prompt: str) -> ToolResult: + """Generate an image from a text prompt. + + Args: + prompt (str): What the image should show. + """ + try: + payload = self._post("/v1/images/generations", self._image_body(prompt)) + assets = [self._download(url) for url in _asset_urls(payload.get("data"))] + return self._image_result(prompt, assets) + except Exception as e: + log_error(f"Failed to generate image using {self.image_model}: {e}") + return ToolResult(content=f"Failed to generate image: {e}") + + async def agenerate_image(self, prompt: str) -> ToolResult: + """Generate an image from a text prompt. + + Args: + prompt (str): What the image should show. + """ + try: + payload = await self._apost("/v1/images/generations", self._image_body(prompt)) + assets = [await self._adownload(url) for url in _asset_urls(payload.get("data"))] + return self._image_result(prompt, assets) + except Exception as e: + log_error(f"Failed to generate image using {self.image_model}: {e}") + return ToolResult(content=f"Failed to generate image: {e}") + + def generate_video(self, prompt: str) -> ToolResult: + """Generate a short video from a text prompt. Takes a minute or more. + + Args: + prompt (str): The scene, subject or action to show. + """ + try: + job = self._post("/v2/video/generations", self._video_body(prompt)) + job_id = job.get("id") + if not isinstance(job_id, str) or not job_id: + return ToolResult(content="Failed to generate video: API did not return a generation id.") + if job.get("status") in _IN_PROGRESS: + job = self._poll( + lambda: self._get("/v2/video/generations", {"generation_id": job_id}), + self.video_poll_interval, + self.video_timeout, + "video generation", + ) + failure = self._job_failure(job, "generate video") + if failure: + return ToolResult(content=failure) + url = _asset_url(job.get("video")) + if url is None: + return ToolResult(content="Failed to generate video: No video data received from API.") + content, mime_type = self._download(url) + return self._video_result(prompt, content, mime_type) + except Exception as e: + log_error(f"Failed to generate video using {self.video_model}: {e}") + return ToolResult(content=f"Failed to generate video: {e}") + + async def agenerate_video(self, prompt: str) -> ToolResult: + """Generate a short video from a text prompt. Takes a minute or more. + + Args: + prompt (str): The scene, subject or action to show. + """ + try: + job = await self._apost("/v2/video/generations", self._video_body(prompt)) + job_id = job.get("id") + if not isinstance(job_id, str) or not job_id: + return ToolResult(content="Failed to generate video: API did not return a generation id.") + if job.get("status") in _IN_PROGRESS: + job = await self._apoll( + lambda: self._aget("/v2/video/generations", {"generation_id": job_id}), + self.video_poll_interval, + self.video_timeout, + "video generation", + ) + failure = self._job_failure(job, "generate video") + if failure: + return ToolResult(content=failure) + url = _asset_url(job.get("video")) + if url is None: + return ToolResult(content="Failed to generate video: No video data received from API.") + content, mime_type = await self._adownload(url) + return self._video_result(prompt, content, mime_type) + except Exception as e: + log_error(f"Failed to generate video using {self.video_model}: {e}") + return ToolResult(content=f"Failed to generate video: {e}") + + def generate_speech(self, text_input: str) -> ToolResult: + """Turn text into spoken audio. + + Args: + text_input (str): The text to read aloud. + """ + try: + payload = self._post("/v1/tts", self._speech_body(text_input)) + url = _asset_url(payload.get("audio")) + if url is None: + return ToolResult(content="Failed to generate speech: No audio data received from API.") + content, mime_type = self._download(url) + return self._speech_result(content, mime_type) + except Exception as e: + log_error(f"Failed to generate speech using {self.speech_model}: {e}") + return ToolResult(content=f"Failed to generate speech: {e}") + + async def agenerate_speech(self, text_input: str) -> ToolResult: + """Turn text into spoken audio. + + Args: + text_input (str): The text to read aloud. + """ + try: + payload = await self._apost("/v1/tts", self._speech_body(text_input)) + url = _asset_url(payload.get("audio")) + if url is None: + return ToolResult(content="Failed to generate speech: No audio data received from API.") + content, mime_type = await self._adownload(url) + return self._speech_result(content, mime_type) + except Exception as e: + log_error(f"Failed to generate speech using {self.speech_model}: {e}") + return ToolResult(content=f"Failed to generate speech: {e}") + + def transcribe_audio(self, audio_path: str) -> str: + """Transcribe an audio file to text. + + Args: + audio_path (str): Path to an audio file inside the toolkit's base directory, or an https URL of one. + """ + try: + if audio_path.startswith(("http://", "https://")): + job = self._post("/v1/stt/create", {**self._transcription_data(), "url": audio_path}) + else: + job = self._post_file("/v1/stt/create", self._transcription_data(), self._local_audio(audio_path)) + job_id = job.get("generation_id") + if not isinstance(job_id, str) or not job_id: + return "Failed to transcribe audio: API did not return a generation id." + if job.get("status") in _IN_PROGRESS: + job = self._poll( + lambda: self._get(f"/v1/stt/{job_id}"), + self.transcription_poll_interval, + self.transcription_timeout, + "transcription", + ) + failure = self._job_failure(job, "transcribe audio") + if failure: + return failure + transcript = self._transcript(job) + if transcript is None: + return "Failed to transcribe audio: No transcript received from API." + log_debug(f"Transcribed {len(transcript)} characters") + return transcript + except Exception as e: + log_error(f"Failed to transcribe audio using {self.transcription_model}: {e}") + return f"Failed to transcribe audio: {e}" + + async def atranscribe_audio(self, audio_path: str) -> str: + """Transcribe an audio file to text. + + Args: + audio_path (str): Path to an audio file inside the toolkit's base directory, or an https URL of one. + """ + try: + if audio_path.startswith(("http://", "https://")): + job = await self._apost("/v1/stt/create", {**self._transcription_data(), "url": audio_path}) + else: + job = await self._apost_file( + "/v1/stt/create", self._transcription_data(), self._local_audio(audio_path) + ) + job_id = job.get("generation_id") + if not isinstance(job_id, str) or not job_id: + return "Failed to transcribe audio: API did not return a generation id." + if job.get("status") in _IN_PROGRESS: + job = await self._apoll( + lambda: self._aget(f"/v1/stt/{job_id}"), + self.transcription_poll_interval, + self.transcription_timeout, + "transcription", + ) + failure = self._job_failure(job, "transcribe audio") + if failure: + return failure + transcript = self._transcript(job) + if transcript is None: + return "Failed to transcribe audio: No transcript received from API." + log_debug(f"Transcribed {len(transcript)} characters") + return transcript + except Exception as e: + log_error(f"Failed to transcribe audio using {self.transcription_model}: {e}") + return f"Failed to transcribe audio: {e}" + + +# --- helpers ------------------------------------------------------------------ + + +def _error_text(error: Any) -> Optional[str]: + """The human-readable part of an error field, whatever shape it took.""" + if isinstance(error, str): + return error or None + if isinstance(error, dict): + for key in ("message", "detail", "name"): + if isinstance(error.get(key), str) and error[key]: + return error[key] + return None + + +def _error_message(response: httpx.Response) -> str: + try: + body = response.json() + except ValueError: + return response.text[:300] + if isinstance(body, dict): + message = body.get("message") + if isinstance(message, str) and message: + return message + nested = _error_text(body.get("error")) + if nested: + return nested + return response.text[:300] + + +def _is_transient(error: Exception) -> bool: + if isinstance(error, httpx.TransportError): + return True + return isinstance(error, AIMLAPIError) and error.status_code in _TRANSIENT_STATUSES + + +def _asset_url(value: Any) -> Optional[str]: + """Generated assets arrive as {"url": ...}, [{"url": ...}] or a bare string.""" + if isinstance(value, list): + value = value[0] if value else None + if isinstance(value, dict): + value = value.get("url") + return value if isinstance(value, str) and value else None + + +def _asset_urls(value: Any) -> List[str]: + if not isinstance(value, list): + return [] + urls = (_asset_url(item) for item in value) + return [url for url in urls if url is not None] + + +def _asset_mime_type(content_type: Optional[str], url: str) -> str: + """The media type of a downloaded asset. + + Signed storage links often answer ``application/octet-stream``; the file + extension in the URL is the next best source. + """ + declared = (content_type or "").split(";")[0].strip().lower() + if declared and declared != "application/octet-stream": + return declared + guessed, _ = mimetypes.guess_type(urlsplit(url).path) + return guessed or declared or "application/octet-stream" + + +def _subtype(mime_type: str) -> str: + """'image/png' -> 'png'; the format field Agno keys media handling on.""" + subtype = mime_type.split("/", 1)[1] if "/" in mime_type else mime_type + return {"mpeg": "mp3", "x-wav": "wav", "quicktime": "mov"}.get(subtype, subtype) diff --git a/libs/agno/agno/tools/neo4j.py b/libs/agno/agno/tools/neo4j.py index 2cd698d9b7f..d15c73bfb05 100644 --- a/libs/agno/agno/tools/neo4j.py +++ b/libs/agno/agno/tools/neo4j.py @@ -27,20 +27,19 @@ def __init__( ): """ Initialize the Neo4jTools toolkit. - Connection parameters (uri/user/password or host/port) can be provided. + Connection parameters (uri/user/password) can be provided. If not provided, falls back to NEO4J_URI, NEO4J_USERNAME, NEO4J_PASSWORD env vars. Args: uri (Optional[str]): The Neo4j URI. user (Optional[str]): The Neo4j username. password (Optional[str]): The Neo4j password. - host (Optional[str]): The Neo4j host. - port (Optional[int]): The Neo4j port. database (Optional[str]): The Neo4j database. - list_labels (bool): Whether to list node labels. - list_relationships (bool): Whether to list relationship types. - get_schema (bool): Whether to get the schema. - run_cypher (bool): Whether to run Cypher queries. + enable_list_labels (bool): Whether to list node labels. + enable_list_relationships (bool): Whether to list relationship types. + enable_get_schema (bool): Whether to get the schema. + enable_run_cypher (bool): Whether to run Cypher queries. + all (bool): Enable all tools. Overrides individual flags when True. Default is False. **kwargs: Additional keyword arguments. """ # Determine the connection URI and credentials diff --git a/libs/agno/agno/tools/python.py b/libs/agno/agno/tools/python.py index 4f1623c4c32..66bb97379aa 100644 --- a/libs/agno/agno/tools/python.py +++ b/libs/agno/agno/tools/python.py @@ -9,10 +9,39 @@ @functools.lru_cache(maxsize=None) def warn() -> None: - logger.warning("PythonTools can run arbitrary code, please provide human supervision.") + logger.warning( + "PythonTools executes arbitrary Python in this process. Provide human supervision and never " + "expose it to untrusted input; safe_globals/safe_locals and restrict_to_base_dir are not a sandbox." + ) class PythonTools(Toolkit): + """Tools for generating, saving, and executing Python code in the current process. + + .. warning:: + ``run_python_code`` and ``save_to_file_and_run`` execute model-generated + Python in this process via ``exec``/``runpy`` with full builtins, imports, + filesystem, and network access. There is no sandbox: an RCE sink if the + agent is prompt-injected. + + ``safe_globals`` / ``safe_locals`` are NOT a security boundary despite the + name: they default to this module's real namespaces and only seed the + execution scope. ``restrict_to_base_dir`` constrains the *path arguments* + of the file helpers (read_file, save_to_file_and_run, ...) but does nothing + to code once it runs: executed code can read ``/etc/passwd``, dump + ``os.environ``, or reach the network regardless of that flag. + + To require human approval before code runs, gate the tools through the + toolkit's confirmation mechanism:: + + PythonTools(requires_confirmation_tools=["run_python_code", "save_to_file_and_run"]) + + To drop the execution tools entirely, use ``exclude_tools=[...]``. For + untrusted input, run code in a real sandbox (separate process or container + with a scrubbed environment, no network, and a read-only mount). See + DaytonaTools for a remote-sandbox alternative. + """ + def __init__( self, base_dir: Optional[Path] = None, @@ -21,10 +50,24 @@ def __init__( restrict_to_base_dir: bool = True, **kwargs, ): + """Initialize PythonTools. + + Args: + base_dir: Root directory for file operations. Defaults to cwd. + safe_globals: Globals namespace seeded into executed code. NOT a + sandbox; defaults to this module's globals. Does not limit what + executed code can import or access. + safe_locals: Locals namespace seeded into executed code. NOT a sandbox; + see safe_globals. + restrict_to_base_dir: If True, confine the *path arguments* of the file + helpers to base_dir. This does not sandbox executed code, which can + still touch any path the process can. Do not rely on it for + untrusted input. + """ self.base_dir: Path = (base_dir or Path.cwd()).resolve() self.restrict_to_base_dir = restrict_to_base_dir - # Restricted global and local scope + # Execution namespaces seeded into exec()/runpy. Not a security boundary. self.safe_globals: dict = safe_globals or globals() self.safe_locals: dict = safe_locals or locals() diff --git a/libs/agno/agno/tools/studio_runner.py b/libs/agno/agno/tools/studio_runner.py index e7eef45a1d4..9baac282ff0 100644 --- a/libs/agno/agno/tools/studio_runner.py +++ b/libs/agno/agno/tools/studio_runner.py @@ -2812,11 +2812,22 @@ async def arun_agent( _agno_agent: Optional[Any] = None, _agno_team: Optional[Any] = None, ) -> str: - """Async variant of run_agent. + """Run an agent and return its result. + + The run executes as the current user and continues that user's + per-conversation session with this agent. A PAUSED status means the run + awaits human approval: the result carries the unresolved requirements + plus the run_id and session_id a continue call must address. A dispatch + refused for a cycle or the depth limit returns an error naming the + lineage; relay it -- do not retry. Args: agent_id (str): Id of the agent to run (a display name or its slug also resolves). message (str): The message to send. + + Returns: + str: JSON object with 'agent_id', 'run_id', 'session_id', 'status', + 'content' and, when paused, 'requirements'. """ # Resolution hits the DB synchronously; keep it off the event loop. actor = getattr(_agno_run_context, "user_id", None) @@ -2868,11 +2879,22 @@ async def arun_team( _agno_agent: Optional[Any] = None, _agno_team: Optional[Any] = None, ) -> str: - """Async variant of run_team. + """Run a team and return its result. + + The run executes as the current user and continues that user's + per-conversation session with this team. A PAUSED status means the run + awaits human approval: the result carries the unresolved requirements + plus the run_id and session_id a continue call must address. A dispatch + refused for a cycle or the depth limit returns an error naming the + lineage; relay it -- do not retry. Args: team_id (str): Id of the team to run (a display name or its slug also resolves). message (str): The message to send. + + Returns: + str: JSON object with 'team_id', 'run_id', 'session_id', 'status', + 'content' and, when paused, 'requirements'. """ actor = getattr(_agno_run_context, "user_id", None) try: @@ -2921,11 +2943,22 @@ async def arun_workflow( _agno_agent: Optional[Any] = None, _agno_team: Optional[Any] = None, ) -> str: - """Async variant of run_workflow. + """Run a workflow and return its final result. + + The run executes as the current user and continues that user's + per-conversation session with this workflow. A PAUSED status means the + run awaits human approval: the result carries the unresolved + requirements plus the run_id and session_id a continue call must address. + A dispatch refused for a cycle or the depth limit returns an error + naming the lineage; relay it -- do not retry. Args: workflow_id (str): Id of the workflow to run (a display name or its slug also resolves). message (str): Input to pass to the first step. + + Returns: + str: JSON object with 'workflow_id', 'run_id', 'session_id', 'status', + 'content' and, when paused, 'requirements'. """ actor = getattr(_agno_run_context, "user_id", None) try: @@ -2969,15 +3002,63 @@ async def arun_workflow( return json.dumps({"error": str(e) or type(e).__name__}) async def alist_agents(self, _agno_run_context: Optional[RunContext] = None) -> str: - """Async variant of list_agents.""" + """List agents this runner can run, newest first. + + Reports the components stored in the platform database, preceded by any + code-defined agents this runner admits (an explicit list, or the + registry under include_all_components). What can be run can be found. + + Returns: + str: JSON object with 'agents' (each {id, name, description}; a row + with status 'draft' has no published version yet, so it will + not dispatch until published), 'count' (returned), 'total' + (every component this runner can run; total > count means the + list is capped -- components beyond the cap still run by + exact id) and 'other_components' (how many runnable teams and + workflows exist -- this list is agents only, so check the + sibling list tools before concluding a component does not + exist). + """ return await asyncio.to_thread(self.list_agents, _agno_run_context=_agno_run_context) async def alist_teams(self, _agno_run_context: Optional[RunContext] = None) -> str: - """Async variant of list_teams.""" + """List teams this runner can run, newest first. + + Reports the components stored in the platform database, preceded by any + code-defined teams this runner admits (an explicit list, or the + registry under include_all_components). What can be run can be found. + + Returns: + str: JSON object with 'teams' (each {id, name, description}; a row + with status 'draft' has no published version yet, so it will + not dispatch until published), 'count' (returned), 'total' + (every component this runner can run; total > count means the + list is capped -- components beyond the cap still run by + exact id) and 'other_components' (how many runnable agents and + workflows exist -- this list is teams only, so check the + sibling list tools before concluding a component does not + exist). + """ return await asyncio.to_thread(self.list_teams, _agno_run_context=_agno_run_context) async def alist_workflows(self, _agno_run_context: Optional[RunContext] = None) -> str: - """Async variant of list_workflows.""" + """List workflows this runner can run, newest first. + + Reports the components stored in the platform database, preceded by any + code-defined workflows this runner admits (an explicit list, or the + registry under include_all_components). What can be run can be found. + + Returns: + str: JSON object with 'workflows' (each {id, name, description}; a row + with status 'draft' has no published version yet, so it will + not dispatch until published), 'count' (returned), 'total' + (every component this runner can run; total > count means the + list is capped -- components beyond the cap still run by + exact id) and 'other_components' (how many runnable agents and + teams exist -- this list is workflows only, so check the + sibling list tools before concluding a component does not + exist). + """ return await asyncio.to_thread(self.list_workflows, _agno_run_context=_agno_run_context) # ------------------------------------------------------------------ diff --git a/libs/agno/agno/tools/superserve.py b/libs/agno/agno/tools/superserve.py index 0d83c6b6d02..68e6848f427 100644 --- a/libs/agno/agno/tools/superserve.py +++ b/libs/agno/agno/tools/superserve.py @@ -597,7 +597,14 @@ def detach_secret(self, agent: Union[Agent, Team], env_key: str) -> str: # Core tools (async) # ------------------------------------------------------------------ async def arun_python_code(self, agent: Union[Agent, Team], code: str) -> str: - """Async variant of run_python_code.""" + """Execute Python code in the sandbox and return its output. + + Args: + code: Python code to execute. + + Returns: + The command output (stdout, stderr, exit code) or an error message. + """ try: sandbox = await self._aget_sandbox(agent) path = f"/tmp/agno_run_{uuid4().hex[:8]}.py" @@ -608,7 +615,14 @@ async def arun_python_code(self, agent: Union[Agent, Team], code: str) -> str: return self._error("Error executing code", e) async def arun_command(self, agent: Union[Agent, Team], command: str) -> str: - """Async variant of run_command.""" + """Execute a shell command in the sandbox. + + Args: + command: Shell command to execute. + + Returns: + The command output (stdout, stderr, exit code) or an error message. + """ try: sandbox = await self._aget_sandbox(agent) result = await sandbox.commands.run(command, timeout_seconds=self.command_timeout) @@ -617,7 +631,15 @@ async def arun_command(self, agent: Union[Agent, Team], command: str) -> str: return self._error("Error executing command", e) async def acreate_file(self, agent: Union[Agent, Team], file_path: str, content: str) -> str: - """Async variant of create_file.""" + """Create or overwrite a file in the sandbox. + + Args: + file_path: Absolute path to the file in the sandbox. + content: Text content to write. + + Returns: + A success message or an error message. + """ try: sandbox = await self._aget_sandbox(agent) await sandbox.files.write(file_path, content) @@ -626,7 +648,14 @@ async def acreate_file(self, agent: Union[Agent, Team], file_path: str, content: return self._error("Error creating file", e) async def aread_file(self, agent: Union[Agent, Team], file_path: str) -> str: - """Async variant of read_file.""" + """Read a file's contents from the sandbox. + + Args: + file_path: Absolute path to the file in the sandbox. + + Returns: + The file contents as text or an error message. + """ try: sandbox = await self._aget_sandbox(agent) return await sandbox.files.read_text(file_path) @@ -634,7 +663,14 @@ async def aread_file(self, agent: Union[Agent, Team], file_path: str) -> str: return self._error("Error reading file", e) async def alist_files(self, agent: Union[Agent, Team], directory: str = "/") -> str: - """Async variant of list_files.""" + """List the contents of a directory in the sandbox. + + Args: + directory: Directory to list (default: root). + + Returns: + The directory listing or an error message. + """ try: sandbox = await self._aget_sandbox(agent) result = await sandbox.commands.run( @@ -647,7 +683,14 @@ async def alist_files(self, agent: Union[Agent, Team], directory: str = "/") -> return self._error("Error listing files", e) async def adelete_file(self, agent: Union[Agent, Team], file_path: str) -> str: - """Async variant of delete_file.""" + """Delete a file or directory in the sandbox. + + Args: + file_path: Absolute path to the file or directory in the sandbox. + + Returns: + A success message or an error message. + """ try: sandbox = await self._aget_sandbox(agent) result = await sandbox.commands.run( @@ -660,7 +703,16 @@ async def adelete_file(self, agent: Union[Agent, Team], file_path: str) -> str: return self._error("Error deleting file", e) async def adownload_directory(self, agent: Union[Agent, Team], sandbox_path: str, local_path: str) -> str: - """Async variant of download_directory.""" + """Download a directory from the sandbox as a zip archive saved locally. + + Args: + sandbox_path: Directory path in the sandbox to download. + local_path: Path within the tool's output directory to write the zip archive to + (e.g. "out.zip"). Must stay inside that directory. + + Returns: + The local path written or an error message. + """ try: sandbox = await self._aget_sandbox(agent) data = await sandbox.files.download_dir(sandbox_path, timeout=self.command_timeout) @@ -672,7 +724,11 @@ async def adownload_directory(self, agent: Union[Agent, Team], sandbox_path: str return self._error("Error downloading directory", e) async def aget_sandbox_info(self, agent: Union[Agent, Team]) -> str: - """Async variant of get_sandbox_info.""" + """Get information about the current sandbox. + + Returns: + JSON with the sandbox id, name, status, and metadata, or an error message. + """ try: sandbox = await self._aget_sandbox(agent) info = await sandbox.get_info() @@ -683,7 +739,11 @@ async def aget_sandbox_info(self, agent: Union[Agent, Team]) -> str: return self._error("Error getting sandbox info", e) async def alist_sandboxes(self) -> str: - """Async variant of list_sandboxes.""" + """List all sandboxes belonging to the team. + + Returns: + JSON list of sandboxes (id, name, status) or an error message. + """ try: sandboxes = await AsyncSandbox.list(api_key=self.api_key, base_url=self.base_url) return json.dumps([{"id": s.id, "name": s.name, "status": s.status.value} for s in sandboxes]) @@ -691,7 +751,11 @@ async def alist_sandboxes(self) -> str: return self._error("Error listing sandboxes", e) async def ashutdown_sandbox(self, agent: Union[Agent, Team]) -> str: - """Async variant of shutdown_sandbox.""" + """Delete the current sandbox and release its resources. + + Returns: + A success message or an error message. + """ try: if self._async_sandbox is None and not self._resolve_sandbox_id(agent): return "No active sandbox to shut down." @@ -706,7 +770,14 @@ async def ashutdown_sandbox(self, agent: Union[Agent, Team]) -> str: return self._error("Error shutting down sandbox", e) async def ashutdown_sandbox_by_id(self, agent: Union[Agent, Team], sandbox_id: str) -> str: - """Async variant of shutdown_sandbox_by_id.""" + """Delete a specific sandbox by its id, e.g. one returned by list_sandboxes. + + Args: + sandbox_id: The id of the sandbox to delete. + + Returns: + A success message or an error message. + """ try: await AsyncSandbox.kill_by_id(sandbox_id, api_key=self.api_key, base_url=self.base_url) if self._is_current_sandbox(agent, sandbox_id): @@ -718,7 +789,14 @@ async def ashutdown_sandbox_by_id(self, agent: Union[Agent, Team], sandbox_id: s return self._error("Error shutting down sandbox", e) async def aget_preview_url(self, agent: Union[Agent, Team], port: int) -> str: - """Async variant of get_preview_url.""" + """Get a public URL for a port exposed inside the sandbox. + + Args: + port: Port a process inside the sandbox is listening on. + + Returns: + A public URL routing to that port, or an error message. + """ try: sandbox = await self._aget_sandbox(agent) return sandbox.get_preview_url(port) @@ -729,7 +807,11 @@ async def aget_preview_url(self, agent: Union[Agent, Team], port: int) -> str: # Lifecycle tools (async, opt-in) # ------------------------------------------------------------------ async def apause_sandbox(self, agent: Union[Agent, Team]) -> str: - """Async variant of pause_sandbox.""" + """Pause the current sandbox to save resources. It can be resumed later. + + Returns: + A success message or an error message. + """ try: sandbox = await self._aget_sandbox(agent) await sandbox.pause() @@ -738,7 +820,11 @@ async def apause_sandbox(self, agent: Union[Agent, Team]) -> str: return self._error("Error pausing sandbox", e) async def aresume_sandbox(self, agent: Union[Agent, Team]) -> str: - """Async variant of resume_sandbox.""" + """Resume the current paused sandbox. + + Returns: + A success message or an error message. + """ try: sandbox = await self._aget_sandbox(agent) await sandbox.resume() @@ -750,7 +836,18 @@ async def aresume_sandbox(self, agent: Union[Agent, Team]) -> str: # Secret tools (async, opt-in) # ------------------------------------------------------------------ async def aattach_secret(self, agent: Union[Agent, Team], env_key: str, secret_name: str) -> str: - """Async variant of attach_secret.""" + """Bind a team secret to the sandbox under an environment variable. + + The sandbox sees a proxy token; the real credential is swapped in only for + outbound requests to the secret's allowed hosts. + + Args: + env_key: Environment variable name the sandbox will see. + secret_name: Name of the team secret to bind. + + Returns: + A success message or an error message. + """ try: sandbox = await self._aget_sandbox(agent) await sandbox.attach_secret(env_key, secret_name) @@ -759,7 +856,14 @@ async def aattach_secret(self, agent: Union[Agent, Team], env_key: str, secret_n return self._error("Error attaching secret", e) async def adetach_secret(self, agent: Union[Agent, Team], env_key: str) -> str: - """Async variant of detach_secret.""" + """Remove a secret binding from the sandbox by its environment variable key. + + Args: + env_key: Environment variable name of the binding to remove. + + Returns: + A success message or an error message. + """ try: sandbox = await self._aget_sandbox(agent) await sandbox.detach_secret(env_key) diff --git a/libs/agno/agno/tools/workflow.py b/libs/agno/agno/tools/workflow.py index d9927be35b0..244de740c53 100644 --- a/libs/agno/agno/tools/workflow.py +++ b/libs/agno/agno/tools/workflow.py @@ -172,8 +172,7 @@ async def async_run_workflow( """Use this tool to execute the workflow with the specified inputs and parameters. After thinking through the requirements, use this tool to run the workflow with appropriate inputs. Args: - input_data: The input data for the workflow (use a `str` for a simple input) - additional_data: The additional data for the workflow. This is a dictionary of key-value pairs that will be passed to the workflow. E.g. {"topic": "food", "style": "Humour"} + input: The input data for the workflow. """ if isinstance(input, dict): input = RunWorkflowInput.model_validate(input) diff --git a/libs/agno/agno/tools/zoom.py b/libs/agno/agno/tools/zoom.py index 609b18bd797..72468c2c132 100644 --- a/libs/agno/agno/tools/zoom.py +++ b/libs/agno/agno/tools/zoom.py @@ -27,7 +27,6 @@ def __init__( client_id (str): The client ID for authentication. If not provided, will use ZOOM_CLIENT_ID env var. client_secret (str): The client secret for authentication. If not provided, will use ZOOM_CLIENT_SECRET env var. timeout (int): Per-request HTTP timeout in seconds. Default is 30. - name (str): The name of the tool. Defaults to "zoom_tool". """ # Get credentials from env vars if not provided self.account_id = account_id or getenv("ZOOM_ACCOUNT_ID") diff --git a/libs/agno/agno/utils/vectors.py b/libs/agno/agno/utils/vectors.py new file mode 100644 index 00000000000..d56075b94e3 --- /dev/null +++ b/libs/agno/agno/utils/vectors.py @@ -0,0 +1,34 @@ +"""Vector math for ranking, without numpy, which is not a core dependency.""" + +from math import sqrt +from typing import List, Sequence + + +def dot(left: Sequence[float], right: Sequence[float]) -> float: + """Dot product, which is cosine similarity when both vectors are unit length.""" + total = 0.0 + for a, b in zip(left, right): + total += a * b + return total + + +def unit(vector: Sequence[float]) -> List[float]: + """Scale to unit length, so repeated similarity checks reduce to a dot product.""" + norm = sqrt(sum(value * value for value in vector)) + if norm <= 0.0: + return [0.0] * len(vector) + return [value / norm for value in vector] + + +def cosine_similarity(left: Sequence[float], right: Sequence[float]) -> float: + """Cosine similarity of two vectors, 0.0 when either has no magnitude.""" + product = 0.0 + left_norm = 0.0 + right_norm = 0.0 + for a, b in zip(left, right): + product += a * b + left_norm += a * a + right_norm += b * b + if left_norm <= 0.0 or right_norm <= 0.0: + return 0.0 + return product / (sqrt(left_norm) * sqrt(right_norm)) diff --git a/libs/agno/agno/vectordb/base.py b/libs/agno/agno/vectordb/base.py index 30b8f655941..626a315b33b 100644 --- a/libs/agno/agno/vectordb/base.py +++ b/libs/agno/agno/vectordb/base.py @@ -1,5 +1,7 @@ from abc import ABC, abstractmethod -from typing import Any, Dict, List, Optional +from contextlib import contextmanager +from contextvars import ContextVar +from typing import Any, Dict, Iterator, List, Optional from agno.exceptions import EmbeddingError from agno.knowledge.document import Document @@ -92,9 +94,42 @@ def is_rate_limit_error(error: BaseException) -> bool: ) +# Set for the duration of one search whose caller applies its own reranker. A ContextVar +# keeps the suspension local to that call, including across awaits and worker threads. +_RERANKER_SUPPRESSED: ContextVar[bool] = ContextVar("agno_reranker_suppressed", default=False) + + +@contextmanager +def suppress_reranker() -> Iterator[None]: + """Hide the store's own reranker from ``VectorDb.reranker`` for this call only.""" + token = _RERANKER_SUPPRESSED.set(True) + try: + yield + finally: + _RERANKER_SUPPRESSED.reset(token) + + class VectorDb(ABC): """Base class for Vector Databases""" + _reranker: Optional[Any] = None + + @property + def reranker(self) -> Optional[Any]: + """The configured reranker, or None while the calling search has suspended it. + + Knowledge suspends it when applying its own reranker. The flag is per-search + rather than an attribute write, so a store shared with another Knowledge, or + used directly, never observes a suspended value from someone else's search. + """ + if _RERANKER_SUPPRESSED.get(): + return None + return self._reranker + + @reranker.setter + def reranker(self, value: Optional[Any]) -> None: + self._reranker = value + def __init__( self, *, diff --git a/libs/agno/agno/vectordb/cassandra/cassandra.py b/libs/agno/agno/vectordb/cassandra/cassandra.py index 03e33605d20..f9d5f28f389 100644 --- a/libs/agno/agno/vectordb/cassandra/cassandra.py +++ b/libs/agno/agno/vectordb/cassandra/cassandra.py @@ -86,6 +86,7 @@ def _row_to_document(self, row: Dict[str, Any]) -> Document: id=row["row_id"], content=row["body_blob"], meta_data=metadata, + embedder=self.embedder, embedding=row["vector"], name=row["document_name"], content_id=metadata.get("content_id"), diff --git a/libs/agno/agno/vectordb/chroma/chromadb.py b/libs/agno/agno/vectordb/chroma/chromadb.py index f8a9a683c24..61c9ae255fd 100644 --- a/libs/agno/agno/vectordb/chroma/chromadb.py +++ b/libs/agno/agno/vectordb/chroma/chromadb.py @@ -1202,6 +1202,7 @@ def fts_search() -> List[Tuple[str, float]]: name=name, meta_data=doc_metadata, content=content, + embedder=self.embedder, embedding=embedding, content_id=content_id, ) @@ -1297,6 +1298,7 @@ def _build_search_results(self, result: QueryResult) -> List[Document]: name=name, meta_data=doc_metadata, content=content, + embedder=self.embedder, embedding=embedding, content_id=content_id, ) @@ -1384,6 +1386,7 @@ def _build_get_results(self, result: Dict[str, Any], query: str = "") -> List[Do name=name, meta_data=doc_metadata, content=content, + embedder=self.embedder, embedding=embedding, content_id=content_id, ) diff --git a/libs/agno/agno/vectordb/couchbase/couchbase.py b/libs/agno/agno/vectordb/couchbase/couchbase.py index dbf6b4a71ef..8cd636373b9 100644 --- a/libs/agno/agno/vectordb/couchbase/couchbase.py +++ b/libs/agno/agno/vectordb/couchbase/couchbase.py @@ -624,6 +624,7 @@ def __get_doc_from_kv(self, response: SearchResult) -> List[Document]: id=doc_id, name=value["name"], content=value["content"], + embedder=self.embedder, meta_data=value["meta_data"], embedding=value["embedding"], content_id=value.get("content_id"), @@ -1413,6 +1414,7 @@ async def __async_get_doc_from_kv(self, response: AsyncSearchIndex) -> List[Docu id=doc_id, name=value.get("name"), content=value.get("content", ""), + embedder=self.embedder, meta_data=value.get("meta_data", {}), embedding=value.get("embedding", []), ) diff --git a/libs/agno/agno/vectordb/elasticsearch/elasticsearch.py b/libs/agno/agno/vectordb/elasticsearch/elasticsearch.py index e5473d9526a..79496c01108 100644 --- a/libs/agno/agno/vectordb/elasticsearch/elasticsearch.py +++ b/libs/agno/agno/vectordb/elasticsearch/elasticsearch.py @@ -988,6 +988,7 @@ def _create_document_from_hit(self, hit: Dict[str, Any]) -> Document: content=doc_data["content"], name=doc_data.get("name"), meta_data=meta_data, + embedder=self.embedder, embedding=doc_data.get("embedding"), usage=doc_data.get("usage"), reranking_score=doc_data.get("reranking_score"), diff --git a/libs/agno/agno/vectordb/opensearch/opensearch.py b/libs/agno/agno/vectordb/opensearch/opensearch.py index 291b5e08547..8e2243d5da8 100644 --- a/libs/agno/agno/vectordb/opensearch/opensearch.py +++ b/libs/agno/agno/vectordb/opensearch/opensearch.py @@ -909,6 +909,7 @@ def _create_document_from_hit(self, hit: Dict[str, Any]) -> Document: content=doc_data["content"], name=doc_data.get("name"), meta_data=meta_data, + embedder=self.embedder, embedding=doc_data.get("embedding"), usage=doc_data.get("usage"), reranking_score=doc_data.get("reranking_score"), diff --git a/libs/agno/agno/vectordb/pineconedb/pineconedb.py b/libs/agno/agno/vectordb/pineconedb/pineconedb.py index 83bf49954c3..d84823537fd 100644 --- a/libs/agno/agno/vectordb/pineconedb/pineconedb.py +++ b/libs/agno/agno/vectordb/pineconedb/pineconedb.py @@ -95,6 +95,7 @@ def __init__( use_hybrid_search: bool = False, hybrid_alpha: float = 0.5, reranker: Optional[Reranker] = None, + return_vectors: bool = False, **kwargs, ): # Validate required parameters @@ -149,6 +150,9 @@ def __init__( log_debug("Embedder not provided, using OpenAIEmbedder as default.") self.embedder: Embedder = _embedder self.reranker: Optional[Reranker] = reranker + # Pinecone omits vectors unless asked. Fetching them enlarges every response, so + # this stays off until a reranker that scores on embeddings needs them. + self.return_vectors: bool = return_vectors @property def client(self) -> Pinecone: @@ -497,6 +501,12 @@ def _hybrid_scale(self, dense: List[float], sparse: Dict[str, Any], alpha: float hdense = [v * alpha for v in dense] return hdense, hsparse + def _include_values(self, include_values: Optional[bool]) -> bool: + """An explicit argument wins; otherwise follow the instance setting.""" + if include_values is not None: + return include_values + return self.return_vectors + def search( self, query: str, @@ -513,7 +523,8 @@ def search( limit (int, optional): The maximum number of results to return. Defaults to 5. filters (Optional[Dict[str, Union[str, float, int, bool, List, dict]]], optional): The filter for the search. Defaults to None. namespace (Optional[str], optional): The namespace to search in. Defaults to None. - include_values (Optional[bool], optional): Whether to include values in the search results. Defaults to None. + include_values (Optional[bool], optional): Whether to include vectors in the results. + Defaults to None, which follows the return_vectors setting on the instance. include_metadata (Optional[bool], optional): Whether to include metadata in the search results. Defaults to None. user_id (Optional[str], optional): Scope results to this user plus shared chunks. Defaults to None, which applies no scope. @@ -543,7 +554,7 @@ def search( top_k=limit, namespace=namespace or self.namespace, filter=filters, - include_values=include_values, + include_values=self._include_values(include_values), include_metadata=True, ) else: @@ -552,7 +563,7 @@ def search( top_k=limit, namespace=namespace or self.namespace, filter=filters, - include_values=include_values, + include_values=self._include_values(include_values), include_metadata=True, ) @@ -560,6 +571,7 @@ def search( Document( content=(result.metadata.get("text", "") if result.metadata is not None else ""), id=result.id, + embedder=self.embedder, embedding=result.values, meta_data=result.metadata, ) diff --git a/libs/agno/agno/vectordb/qdrant/qdrant.py b/libs/agno/agno/vectordb/qdrant/qdrant.py index 3981ebe8453..b0bb96e9911 100644 --- a/libs/agno/agno/vectordb/qdrant/qdrant.py +++ b/libs/agno/agno/vectordb/qdrant/qdrant.py @@ -600,6 +600,13 @@ async def async_upsert( await asyncio.to_thread(self._delete_by_content_hash, content_hash, user_id) await self.async_insert(content_hash=content_hash, documents=documents, filters=filters, user_id=user_id) + def _dense_vector(self, vector: Any) -> Optional[List[float]]: + """Named-vector searches return a mapping, so pull the dense vector out of it.""" + if isinstance(vector, dict): + dense = vector.get(self.dense_vector_name) + return list(dense) if dense is not None else None + return vector + def search( self, query: str, @@ -821,7 +828,7 @@ def _build_search_results(self, results, query: str) -> List[Document]: meta_data=result.payload["meta_data"], content=result.payload["content"], embedder=self.embedder, - embedding=result.vector, # type: ignore + embedding=self._dense_vector(result.vector), usage=result.payload.get("usage"), content_id=result.payload.get("content_id"), ) diff --git a/libs/agno/agno/vectordb/upstashdb/upstashdb.py b/libs/agno/agno/vectordb/upstashdb/upstashdb.py index 3ee73397366..e94213f49eb 100644 --- a/libs/agno/agno/vectordb/upstashdb/upstashdb.py +++ b/libs/agno/agno/vectordb/upstashdb/upstashdb.py @@ -482,6 +482,7 @@ def search( content=result.data, id=result.id, meta_data=result.metadata or {}, + embedder=self.embedder, embedding=result.vector, ) ) diff --git a/libs/agno/agno/workflow/parallel.py b/libs/agno/agno/workflow/parallel.py index c2c4ccf9aa1..18d2bfe614b 100644 --- a/libs/agno/agno/workflow/parallel.py +++ b/libs/agno/agno/workflow/parallel.py @@ -286,7 +286,7 @@ def _build_aggregated_content(self, step_outputs: List[StepOutput]) -> str: for i, output in enumerate(step_outputs): step_name = output.step_name or f"Step {i + 1}" - content = output.content or "" + content = output.content # Add status indicator if output.success is False: @@ -295,7 +295,7 @@ def _build_aggregated_content(self, step_outputs: List[StepOutput]) -> str: status_icon = "✅ SUCCESS:" aggregated += f"### {status_icon} {step_name}\n" - if content and str(content).strip(): + if content is not None and str(content).strip(): aggregated += f"{content}\n\n" else: aggregated += "*(No content)*\n\n" diff --git a/libs/agno/agno/workflow/step.py b/libs/agno/agno/workflow/step.py index 6e79fd6ee5c..2fa20244a42 100644 --- a/libs/agno/agno/workflow/step.py +++ b/libs/agno/agno/workflow/step.py @@ -2527,7 +2527,9 @@ def _store_executor_response( if isinstance(member_response, RunOutput): workflow_run_response.step_executor_runs.append(member_response) - def _get_deepest_content_from_step_output(self, step_output: "StepOutput") -> Optional[str]: + def _get_deepest_content_from_step_output( + self, step_output: "StepOutput" + ) -> Optional[Union[str, Dict[str, Any], List[Any], BaseModel]]: """ Extract the deepest content from a step output, handling nested structures like Steps, Router, Loop, etc. @@ -2543,16 +2545,16 @@ def _get_deepest_content_from_step_output(self, step_output: "StepOutput") -> Op aggregated_parts = [] for i, inner_step in enumerate(step_output.steps): inner_content = self._get_deepest_content_from_step_output(inner_step) - if inner_content: + if inner_content is not None and str(inner_content).strip(): step_name = inner_step.step_name or f"Step {i + 1}" aggregated_parts.append(f"=== {step_name} ===\n{inner_content}") - return "\n\n".join(aggregated_parts) if aggregated_parts else step_output.content # type: ignore + return "\n\n".join(aggregated_parts) if aggregated_parts else step_output.content # For other nested step types, recursively get content from the last nested step return self._get_deepest_content_from_step_output(step_output.steps[-1]) # For regular steps, return their content - return step_output.content # type: ignore + return step_output.content def _prepare_message( self, diff --git a/libs/agno/agno/workflow/types.py b/libs/agno/agno/workflow/types.py index d5023f094af..8f6217c2c7b 100644 --- a/libs/agno/agno/workflow/types.py +++ b/libs/agno/agno/workflow/types.py @@ -389,12 +389,16 @@ def get_step_content(self, step_name: str) -> Optional[Union[str, Dict[str, str] # Return dict with {step_name: content} for each sub-step parallel_content = {} for sub_step in step_output.steps: - if sub_step.step_name and sub_step.content: + if sub_step.step_name and sub_step.content is not None and str(sub_step.content).strip(): # Check if this sub-step has its own nested steps (like Condition -> Research Step) if sub_step.steps and len(sub_step.steps) > 0: # This is a composite step (like Condition) - get content from its nested steps for nested_step in sub_step.steps: - if nested_step.step_name and nested_step.content: + if ( + nested_step.step_name + and nested_step.content is not None + and str(nested_step.content).strip() + ): parallel_content[nested_step.step_name] = str(nested_step.content) else: # This is a direct step - use its content @@ -425,7 +429,7 @@ def get_all_previous_content(self) -> str: content_parts = [] for step_name, output in self.previous_step_outputs.items(): - if output.content: + if output.content is not None and str(output.content).strip(): content_parts.append(f"=== {step_name} ===\n{output.content}") return "\n\n".join(content_parts) diff --git a/libs/agno/pyproject.toml b/libs/agno/pyproject.toml index aaca1ee0175..1eb4730ed1c 100644 --- a/libs/agno/pyproject.toml +++ b/libs/agno/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "agno" -version = "3.0.9" +version = "3.0.10" description = "The programming language for agentic software." requires-python = ">=3.9,<4" readme = "README.md" diff --git a/libs/agno/tests/integration/agent/test_tool_hooks.py b/libs/agno/tests/integration/agent/test_tool_hooks.py index 1e5ca382ad9..5b843ff20f9 100644 --- a/libs/agno/tests/integration/agent/test_tool_hooks.py +++ b/libs/agno/tests/integration/agent/test_tool_hooks.py @@ -238,7 +238,11 @@ def test_pre_post_hook_receives_messages(): def test_tool_hook_receives_messages(): """Test that tool hooks receive run messages via run_context.messages.""" captured_messages.clear() - agent = Agent(tools=[modulo], tool_hooks=[messages_tool_hook]) + agent = Agent( + tools=[modulo], + tool_hooks=[messages_tool_hook], + instructions="Always use the modulo tool to compute remainders.", + ) response: RunOutput = agent.run("Compute 10 mod 3") diff --git a/libs/agno/tests/unit/app/test_agui_app.py b/libs/agno/tests/unit/app/test_agui_app.py index f1f5e1c1fac..cd8048e68b8 100644 --- a/libs/agno/tests/unit/app/test_agui_app.py +++ b/libs/agno/tests/unit/app/test_agui_app.py @@ -14,6 +14,7 @@ UserMessage, VideoInputContent, ) +from pydantic import ValidationError from agno.models.response import ToolExecution from agno.os.interfaces.agui.input import extract_context, extract_media, extract_user_input @@ -1868,7 +1869,8 @@ async def mock_stream(): # Verify the delta contains the right operations delta_event = events[delta_idx] - delta_paths = [op["path"] for op in delta_event.delta] + # ag-ui-protocol 1.0 parses each JSON Patch entry into a typed operation; earlier releases keep the dict. + delta_paths = [op["path"] if isinstance(op, dict) else op.path for op in delta_event.delta] assert "/counter" in delta_paths assert "/status" in delta_paths @@ -2027,6 +2029,11 @@ def test_extract_media_all_types(): def test_extract_media_binary_content(): """Test AG-UI binary content is routed to the matching Agno media bucket.""" + try: + UserMessage(id="probe", content=[BinaryInputContent(mime_type="image/png", data="aGk=")]) + except ValidationError: + pytest.skip("ag-ui-protocol 1.0 removed the binary content part, so no message can carry one") + image_bytes = b"binary-image" audio_bytes = b"binary-audio" video_bytes = b"binary-video" diff --git a/libs/agno/tests/unit/knowledge/chunking/test_fixed_size_chunking.py b/libs/agno/tests/unit/knowledge/chunking/test_fixed_size_chunking.py index edbca7e1abb..aa80ea37220 100644 --- a/libs/agno/tests/unit/knowledge/chunking/test_fixed_size_chunking.py +++ b/libs/agno/tests/unit/knowledge/chunking/test_fixed_size_chunking.py @@ -32,3 +32,14 @@ def test_long_document_still_chunks_with_overlap_and_no_duplication(): assert len(chunks) == 7 assert [len(c.content) for c in chunks] == [20, 20, 20, 20, 20, 20, 10] + + +def test_repeated_whitespace_collapses_without_flattening_lines(): + """Test that cleaning keeps newlines and tabs while collapsing repeated whitespace.""" + strategy = FixedSizeChunking(chunk_size=100) + doc = Document(name="structured", content="Steps:\n\n\n- Open a ticket\n\t- Attach the receipt") + + chunks = strategy.chunk(doc) + + assert len(chunks) == 1 + assert chunks[0].content == "Steps:\n- Open a ticket\n\t- Attach the receipt" diff --git a/libs/agno/tests/unit/knowledge/test_knowledge_reranker.py b/libs/agno/tests/unit/knowledge/test_knowledge_reranker.py new file mode 100644 index 00000000000..f6eb92a7a94 --- /dev/null +++ b/libs/agno/tests/unit/knowledge/test_knowledge_reranker.py @@ -0,0 +1,504 @@ +"""Knowledge-level reranking: over-fetch, trimming and failure handling.""" + +from typing import Dict, List, Optional + +import pytest +from pydantic import Field + +from agno.knowledge.document import Document +from agno.knowledge.knowledge import Knowledge +from agno.knowledge.reranker.base import Reranker + + +class StubVectorDb: + """Records the limit it was asked for and returns that many documents.""" + + def __init__(self, available: int = 100): + self.available = available + self.requested_limit: Optional[int] = None + + def exists(self) -> bool: + return True + + def create(self) -> None: # pragma: no cover - exists() is always True + raise AssertionError("create should not be called") + + def search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + self.requested_limit = limit + count = min(limit, self.available) + return [Document(id=str(i), content=f"doc {i}") for i in range(count)] + + async def async_search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + return self.search(query=query, limit=limit, filters=filters) + + +class NoAsyncVectorDb(StubVectorDb): + """Exercises the asearch fallback for adapters without async support.""" + + async def async_search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + raise NotImplementedError + + +class ReverseReranker(Reranker): + """Reverses order so the effect of reranking is observable. + + Widens like a selecting reranker, so the pool behaviour is exercised. + """ + + candidate_multiplier: int = Field(default=5, ge=1) + + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + return list(reversed(documents)) + + +class FailingReranker(Reranker): + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + raise RuntimeError("reranker unavailable") + + +def test_search_without_reranker_is_unchanged(): + db = StubVectorDb() + knowledge = Knowledge(vector_db=db) + + results = knowledge.search("q", max_results=5) + + assert db.requested_limit == 5 + assert len(results) == 5 + + +def test_search_over_fetches_when_reranker_is_set(): + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + results = knowledge.search("q", max_results=5) + + assert db.requested_limit == 25 + assert len(results) == 5 + # Reversing 25 candidates surfaces the tail, which plain search would never return. + assert results[0].id == "24" + + +def test_over_fetch_is_capped(): + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker(max_candidates=30)) + + knowledge.search("q", max_results=10) + + assert db.requested_limit == 30 + + +def test_reranker_failure_falls_back_to_vector_db_order(): + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=FailingReranker()) + + results = knowledge.search("q", max_results=5) + + assert len(results) == 5 + assert results[0].id == "0" + + +def test_fewer_candidates_than_requested_is_not_padded(): + db = StubVectorDb(available=3) + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + results = knowledge.search("q", max_results=5) + + assert len(results) == 3 + + +@pytest.mark.asyncio +async def test_asearch_applies_reranker(): + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + results = await knowledge.asearch("q", max_results=5) + + assert db.requested_limit == 25 + assert results[0].id == "24" + + +@pytest.mark.asyncio +async def test_asearch_applies_reranker_on_sync_fallback(): + db = NoAsyncVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + results = await knowledge.asearch("q", max_results=5) + + assert db.requested_limit == 25 + assert results[0].id == "24" + + +@pytest.mark.parametrize("kwargs", [{"candidate_multiplier": 0}, {"candidate_multiplier": -1}, {"max_candidates": 0}]) +def test_invalid_pool_configuration_is_rejected(kwargs): + from pydantic import ValidationError + + with pytest.raises(ValidationError): + ReverseReranker(**kwargs) + + +class ValueErrorReranker(Reranker): + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + raise ValueError("misconfigured") + + +def test_reranker_value_error_propagates(): + # Misconfiguration must surface rather than degrade to unreranked results. + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ValueErrorReranker()) + + with pytest.raises(ValueError, match="misconfigured"): + knowledge.search("q", max_results=5) + + +@pytest.mark.asyncio +async def test_reranker_value_error_propagates_async(): + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ValueErrorReranker()) + + with pytest.raises(ValueError, match="misconfigured"): + await knowledge.asearch("q", max_results=5) + + +def test_search_limit_never_drops_below_requested_results(): + # The ceiling caps the widening, not the caller's own request. + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ReverseReranker(max_candidates=100)) + + assert knowledge._search_limit(150) == 150 + + +def test_over_fetch_is_capped_between_requested_and_ceiling(): + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ReverseReranker(max_candidates=100)) + + assert knowledge._search_limit(10) == 50 + assert knowledge._search_limit(30) == 100 + + +def test_large_max_results_returns_everything_requested(): + db = StubVectorDb(available=200) + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker(max_candidates=100)) + + results = knowledge.search("q", max_results=150) + + assert db.requested_limit == 150 + assert len(results) == 150 + + +@pytest.mark.asyncio +async def test_async_rerank_does_not_block_the_event_loop(): + import asyncio + import threading + + main_thread = threading.get_ident() + seen: List[int] = [] + + class ThreadRecordingReranker(Reranker): + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + seen.append(threading.get_ident()) + return documents + + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ThreadRecordingReranker()) + await knowledge.asearch("q", max_results=5) + + assert seen and seen[0] != main_thread + assert asyncio.get_running_loop().is_running() + + +def test_page_store_results_are_reranked(): + # Page-backed knowledge returns before the vector db, so it needs its own wiring. + from agno.knowledge.page import SearchResult + + class StubPageStore: + pass + + recorded: Dict[str, int] = {} + + knowledge = Knowledge.__new__(Knowledge) + knowledge.page_store = StubPageStore() + knowledge.max_results = 10 + knowledge.reranker = ReverseReranker() + + def fake_search_pages(query, *, limit=10, **kwargs): + recorded["limit"] = limit + return SearchResult(results=[], partial=False) + + knowledge.search_pages = fake_search_pages # type: ignore[method-assign] + knowledge._page_documents = staticmethod( # type: ignore[method-assign] + lambda result: [Document(id=str(i), content=f"doc {i}") for i in range(25)] + ) + + results = knowledge.search("q", max_results=5) + + # Clamped to the page search ceiling rather than the full 5x widening. + assert recorded["limit"] == 20 + assert len(results) == 5 + assert results[0].id == "24" + + +@pytest.mark.parametrize("max_results", [5, 10, 20]) +def test_page_search_limit_stays_within_the_coordinator_ceiling(max_results): + # PageCoordinator.search raises invalid_search_query outside 1..20, so the widened + # page fetch has to clamp rather than pass a multiplied limit straight through. + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ReverseReranker()) + + assert 1 <= knowledge._page_search_limit(max_results) <= 20 + + +def test_page_search_limit_still_widens_when_it_fits(): + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=ReverseReranker()) + + assert knowledge._page_search_limit(2) == 10 + + +def test_reranker_with_the_older_signature_still_works(): + # Rerankers written before `limit` was added, including ones outside this repo, + # must keep working rather than raising an unexpected-keyword TypeError. + class LegacyReranker(Reranker): + candidate_multiplier: int = Field(default=5, ge=1) + + def rerank(self, query: str, documents: List[Document]) -> List[Document]: + return list(reversed(documents)) + + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=LegacyReranker()) + + results = knowledge.search("q", max_results=5) + + assert len(results) == 5 + assert results[0].id == "24" + + +def test_limit_is_passed_to_rerankers_that_accept_it(): + seen = {} + + class LimitAwareReranker(Reranker): + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + seen["limit"] = limit + return documents + + knowledge = Knowledge(vector_db=StubVectorDb(), reranker=LimitAwareReranker()) + knowledge.search("q", max_results=5) + + assert seen["limit"] == 5 + + +def test_a_reranker_can_opt_out_of_widening(): + # The base default widens, since reranking earns its cost by rescuing documents + # ranked below the cutoff. A reranker that only reorders can opt down to 1. + class ReorderOnlyReranker(Reranker): + candidate_multiplier: int = Field(default=1, ge=1) + + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + return documents + + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ReorderOnlyReranker()) + + knowledge.search("q", max_results=5) + + assert db.requested_limit == 5 + + +def test_the_base_default_widens_the_fetch(): + class ScoringReranker(Reranker): + def rerank(self, query: str, documents: List[Document], limit: Optional[int] = None) -> List[Document]: + return documents + + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ScoringReranker()) + + knowledge.search("q", max_results=5) + + assert db.requested_limit == 15 + + +def test_pool_size_is_configured_on_the_reranker(): + db = StubVectorDb() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker(candidate_multiplier=3)) + + knowledge.search("q", max_results=5) + + assert db.requested_limit == 15 + + +class RerankerAwareVectorDb(StubVectorDb): + """Runs its own reranker the way the adapters do, and records that it ran. + + Reads the reranker through VectorDb's property, which is what hides it from a + search that has suspended it. + """ + + def __init__(self, available: int = 100): + super().__init__(available=available) + self._reranker: Optional[Reranker] = None + self.reranker_ran = False + + @property + def reranker(self) -> Optional[Reranker]: + from agno.vectordb.base import VectorDb + + return VectorDb.reranker.fget(self) # type: ignore[attr-defined] + + @reranker.setter + def reranker(self, value: Optional[Reranker]) -> None: + self._reranker = value + + def search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + documents = super().search(query=query, limit=limit, filters=filters) + reranker = self.reranker + if reranker is not None: + self.reranker_ran = True + documents = reranker.rerank(query=query, documents=documents) + return documents + + +def test_knowledge_reranker_wins_over_the_vector_db_one(): + db = RerankerAwareVectorDb() + db.reranker = ReverseReranker() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + results = knowledge.search("q", max_results=5) + + assert db.reranker_ran is False + # Reversed once, by Knowledge, rather than twice. + assert results[0].id == "24" + + +def test_the_vector_db_reranker_is_restored_after_the_search(): + db = RerankerAwareVectorDb() + original = ReverseReranker() + db.reranker = original + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + knowledge.search("q", max_results=5) + + assert db.reranker is original + + +def test_the_vector_db_reranker_still_runs_when_knowledge_has_none(): + db = RerankerAwareVectorDb() + db.reranker = ReverseReranker() + knowledge = Knowledge(vector_db=db) + + knowledge.search("q", max_results=5) + + assert db.reranker_ran is True + + +@pytest.mark.asyncio +async def test_knowledge_reranker_wins_in_async_search(): + db = RerankerAwareVectorDb() + db.reranker = ReverseReranker() + knowledge = Knowledge(vector_db=db, reranker=ReverseReranker()) + + await knowledge.asearch("q", max_results=5) + + assert db.reranker_ran is False + assert db.reranker is not None + + +def _captured_warnings(monkeypatch) -> List[str]: + """Agno's logger sets propagate=False, so caplog never sees these.""" + import agno.knowledge.knowledge as knowledge_module + + messages: List[str] = [] + monkeypatch.setattr(knowledge_module, "log_warning", lambda message, *a, **k: messages.append(str(message))) + return messages + + +def test_configuring_both_rerankers_warns(monkeypatch): + messages = _captured_warnings(monkeypatch) + db = RerankerAwareVectorDb() + db.reranker = ReverseReranker() + + Knowledge(vector_db=db, reranker=ReverseReranker()) + + assert any("set on both Knowledge and the vector db" in message for message in messages) + + +def test_configuring_one_reranker_does_not_warn(monkeypatch): + messages = _captured_warnings(monkeypatch) + db = RerankerAwareVectorDb() + + Knowledge(vector_db=db, reranker=ReverseReranker()) + + assert not any("set on both Knowledge and the vector db" in message for message in messages) + + +@pytest.mark.asyncio +async def test_a_shipped_reranker_with_the_older_signature_survives_arerank(): + # CohereReranker does not override arerank, so the base one must not forward a + # limit its two-argument rerank cannot accept. + cohere = pytest.importorskip("agno.knowledge.reranker.cohere") + + reranker = cohere.CohereReranker(api_key="test") + assert reranker.accepts_limit() is False + + # Reaches rerank without a TypeError; the empty list short-circuits the API call. + assert await reranker.arerank(query="q", documents=[], limit=5) == [] + + +@pytest.mark.asyncio +async def test_suspension_does_not_leak_to_another_knowledge_sharing_the_store(): + # One vector db behind two Knowledge instances is a normal setup. Suspending the + # store's reranker for one search must not hide it from the other. The barrier makes + # the overlap deterministic: B reads the attribute while A's window is open. + import asyncio + + observed: Dict[str, Optional[str]] = {} + inside_a = asyncio.Event() + b_has_read = asyncio.Event() + + class ObservingVectorDb(RerankerAwareVectorDb): + async def async_search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + if query == "from-a": + inside_a.set() + await asyncio.wait_for(b_has_read.wait(), timeout=5) + else: + await asyncio.wait_for(inside_a.wait(), timeout=5) + observed[query] = "set" if self.reranker is not None else None + if query == "from-b": + b_has_read.set() + return StubVectorDb.search(self, query=query, limit=limit, filters=filters) + + db = ObservingVectorDb() + db.reranker = ReverseReranker() + with_own = Knowledge(vector_db=db, reranker=ReverseReranker()) + relies_on_db = Knowledge(vector_db=db) + + await asyncio.gather( + with_own.asearch("from-a", max_results=3), + relies_on_db.asearch("from-b", max_results=3), + ) + + # A suppressed it for itself; B, reading inside that window, still sees its own. + assert observed["from-a"] is None + assert observed["from-b"] == "set" + + +def test_no_reranker_returns_the_adapter_result_untouched(): + # The default path every current user is on: without a reranker the adapter's list + # is returned as produced, not re-sliced by Knowledge. + class OverReturningVectorDb(StubVectorDb): + def search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + self.requested_limit = limit + # An adapter that hands back more than asked for must not be silently trimmed. + return [Document(id=str(i), content=f"doc {i}") for i in range(limit + 2)] + + db = OverReturningVectorDb() + knowledge = Knowledge(vector_db=db) + + results = knowledge.search("q", max_results=5) + + assert db.requested_limit == 5 + assert len(results) == 7 + + +@pytest.mark.asyncio +async def test_no_reranker_returns_the_adapter_result_untouched_async(): + class OverReturningVectorDb(StubVectorDb): + async def async_search(self, query: str, limit: int = 5, filters=None) -> List[Document]: + self.requested_limit = limit + return [Document(id=str(i), content=f"doc {i}") for i in range(limit + 2)] + + knowledge = Knowledge(vector_db=OverReturningVectorDb()) + + results = await knowledge.asearch("q", max_results=5) + + assert len(results) == 7 diff --git a/libs/agno/tests/unit/knowledge/test_mmr_reranker.py b/libs/agno/tests/unit/knowledge/test_mmr_reranker.py new file mode 100644 index 00000000000..87796c83050 --- /dev/null +++ b/libs/agno/tests/unit/knowledge/test_mmr_reranker.py @@ -0,0 +1,295 @@ +"""MMR reranking: diversity selection, configuration and embedding requirements.""" + +from typing import List, Optional + +import pytest + +from agno.knowledge.document import Document +from agno.knowledge.reranker.mmr import MMRReranker +from agno.utils.vectors import cosine_similarity + + +class StubEmbedder: + """Returns a fixed query embedding without calling a provider.""" + + def __init__(self, embedding: Optional[List[float]] = None): + self.embedding = [1.0, 0.0] if embedding is None else embedding + + def get_embedding(self, text: str) -> List[float]: + return self.embedding + + async def async_get_embedding(self, text: str) -> List[float]: + return self.embedding + + +def _documents() -> List[Document]: + """Two near-duplicates, then a document that is less relevant but far from them. + + Relevance alone ranks a > b > c. Selecting "a" first makes "b" redundant, so an + even relevance/diversity split prefers "c" despite its lower relevance. + """ + embedder = StubEmbedder() + return [ + Document(id="a", content="a", embedding=[1.0, 0.5], embedder=embedder), + Document(id="b", content="b", embedding=[1.0, 0.55], embedder=embedder), + Document(id="c", content="c", embedding=[1.0, -0.7], embedder=embedder), + ] + + +def test_cosine_similarity_of_identical_vectors_is_one(): + assert cosine_similarity([1.0, 0.0], [1.0, 0.0]) == pytest.approx(1.0) + + +def test_cosine_similarity_of_orthogonal_vectors_is_zero(): + assert cosine_similarity([1.0, 0.0], [0.0, 1.0]) == pytest.approx(0.0) + + +def test_zero_vector_does_not_divide_by_zero(): + assert cosine_similarity([0.0, 0.0], [1.0, 0.0]) == 0.0 + + +def test_diversity_beats_the_near_duplicate(): + results = MMRReranker(lambda_mult=0.5, top_n=2).rerank("q", _documents()) + + # "b" is the closer match, but it nearly duplicates "a", so the distinct doc wins. + assert [doc.id for doc in results] == ["a", "c"] + + +def test_pure_relevance_keeps_the_near_duplicate(): + results = MMRReranker(lambda_mult=1.0, top_n=2).rerank("q", _documents()) + + assert [doc.id for doc in results] == ["a", "b"] + + +def test_reranking_score_is_the_score_at_selection_time(): + # MMR scores are not descending: the pool shrinks as redundancy grows, so a later + # pick can score above an earlier one. List order, not score order, is the result. + embedder = StubEmbedder() + documents = [ + Document(id="a", content="a", embedding=[-1.0, 0.0], embedder=embedder), + Document(id="b", content="b", embedding=[-0.9, 0.44], embedder=embedder), + Document(id="c", content="c", embedding=[-0.9, -0.44], embedder=embedder), + ] + + results = MMRReranker(lambda_mult=0.5).rerank("q", documents) + + scores = [doc.reranking_score for doc in results] + assert scores != sorted(scores, reverse=True) + assert [doc.id for doc in results] == ["b", "c", "a"] + + +def test_top_n_defaults_to_all_documents(): + results = MMRReranker().rerank("q", _documents()) + + assert len(results) == 3 + + +def test_top_n_larger_than_input_is_clamped(): + results = MMRReranker(top_n=10).rerank("q", _documents()) + + assert len(results) == 3 + + +def test_empty_documents_returns_empty(): + assert MMRReranker().rerank("q", []) == [] + + +def test_missing_embeddings_raise_rather_than_silently_passing_through(): + documents = _documents() + documents[1].embedding = None + + with pytest.raises(ValueError, match="requires embeddings"): + MMRReranker().rerank("q", documents) + + +def test_missing_embedder_raises(): + documents = _documents() + for document in documents: + document.embedder = None + + with pytest.raises(ValueError, match="needs an embedder"): + MMRReranker().rerank("q", documents) + + +def test_embedder_is_taken_from_any_result_that_has_one(): + documents = _documents() + documents[0].embedder = None + + results = MMRReranker(lambda_mult=0.5, top_n=2).rerank("q", documents) + + assert [doc.id for doc in results] == ["a", "c"] + + +@pytest.mark.asyncio +async def test_arerank_matches_sync_selection(): + results = await MMRReranker(lambda_mult=0.5, top_n=2).arerank("q", _documents()) + + assert [doc.id for doc in results] == ["a", "c"] + + +def test_mismatched_embedding_dimensions_raise(): + documents = _documents() + documents[1].embedding = [1.0, 0.5, 0.25] + + with pytest.raises(ValueError, match="one dimension"): + MMRReranker().rerank("q", documents) + + +def test_query_embedding_of_wrong_dimension_raises(): + documents = _documents() + for document in documents: + document.embedder = StubEmbedder(embedding=[1.0, 0.0, 0.0]) + + with pytest.raises(ValueError, match="one dimension"): + MMRReranker().rerank("q", documents) + + +def test_embedder_returning_no_vector_raises(): + documents = _documents() + for document in documents: + document.embedder = StubEmbedder(embedding=[]) + + with pytest.raises(ValueError, match="could not embed the query"): + MMRReranker().rerank("q", documents) + + +def test_numpy_embeddings_are_supported(): + # PgVector returns embeddings as numpy arrays, whose truth value is ambiguous. + numpy = pytest.importorskip("numpy") + + embedder = StubEmbedder(embedding=numpy.array([1.0, 0.0])) + documents = [ + Document(id="a", content="a", embedding=numpy.array([1.0, 0.5]), embedder=embedder), + Document(id="b", content="b", embedding=numpy.array([1.0, 0.55]), embedder=embedder), + Document(id="c", content="c", embedding=numpy.array([1.0, -0.7]), embedder=embedder), + ] + + results = MMRReranker(lambda_mult=0.5, top_n=2).rerank("q", documents) + + assert [doc.id for doc in results] == ["a", "c"] + + +@pytest.mark.parametrize("kwargs", [{"lambda_mult": True}, {"top_n": True}]) +def test_booleans_are_rejected_as_numeric_config(kwargs): + # bool is an int subclass, so True would otherwise coerce to 1.0 / 1. + from pydantic import ValidationError + + with pytest.raises(ValidationError): + MMRReranker(**kwargs) + + +def test_named_vector_mapping_is_not_treated_as_an_embedding(): + # Qdrant hybrid/keyword searches return a dict of named vectors. + documents = _documents() + documents[1].embedding = {"dense": [1.0, 0.55]} # type: ignore[assignment] + + with pytest.raises(ValueError, match="one dimension"): + MMRReranker().rerank("q", documents) + + +def test_reranking_score_is_the_mmr_score_not_a_rank_ordinal(): + results = MMRReranker(lambda_mult=0.5, top_n=2).rerank("q", _documents()) + + # Rank ordinals would be 2.0 and 1.0; real MMR scores are bounded by lambda_mult. + scores = [doc.reranking_score for doc in results] + assert all(score <= 1.0 for score in scores) + + +def test_input_documents_are_not_scored_in_place(): + documents = _documents() + + MMRReranker(lambda_mult=0.5, top_n=2).rerank("q", documents) + + assert all(document.reranking_score is None for document in documents) + + +def test_zero_vector_is_rejected_as_a_broken_embedding(): + documents = _documents() + documents[1].embedding = [0.0, 0.0] + + with pytest.raises(ValueError, match="requires embeddings"): + MMRReranker().rerank("q", documents) + + +@pytest.mark.parametrize("kwargs", [{"lambda_mult": -0.1}, {"lambda_mult": 1.1}, {"top_n": 0}, {"top_n": -1}]) +def test_out_of_range_config_is_rejected_at_construction(kwargs): + from pydantic import ValidationError + + with pytest.raises(ValidationError): + MMRReranker(**kwargs) + + +def test_pure_diversity_does_not_depend_on_input_order(): + # At lambda_mult=0.0 every candidate ties on the first pick, so the seed cannot be + # scored with the MMR formula or the input order decides the whole selection. + import itertools + + embeddings = {"a": [1.0, 0.1], "b": [1.0, 0.5], "c": [1.0, -0.9]} + + def documents(order): + embedder = StubEmbedder() + return [Document(id=key, content=key, embedding=embeddings[key], embedder=embedder) for key in order] + + selections = { + tuple(doc.id for doc in MMRReranker(lambda_mult=0.0, top_n=2).rerank("q", documents(list(order)))) + for order in itertools.permutations("abc") + } + + assert len(selections) == 1 + + +def test_selection_matches_the_unoptimised_formula(): + # Pins the running-redundancy optimisation to the definition it replaced. + import random + + from agno.utils.vectors import cosine_similarity + + def reference(query_embedding, documents, lambda_mult, limit): + embeddings = [doc.embedding for doc in documents] + relevance = [cosine_similarity(query_embedding, embedding) for embedding in embeddings] + remaining = list(range(len(documents))) + first = max(remaining, key=lambda candidate: relevance[candidate]) + selected = [first] + remaining.remove(first) + while remaining and len(selected) < limit: + best_index, best_score = remaining[0], float("-inf") + for candidate in remaining: + redundancy = max(cosine_similarity(embeddings[candidate], embeddings[j]) for j in selected) + score = lambda_mult * relevance[candidate] - (1.0 - lambda_mult) * redundancy + if score > best_score: + best_score, best_index = score, candidate + selected.append(best_index) + remaining.remove(best_index) + return [documents[i].id for i in selected] + + query_embedding = [1.0] + [0.0] * 15 + + def build(seed): + rnd = random.Random(seed) + embedder = StubEmbedder(embedding=query_embedding) + return [ + Document( + id=str(i), + content=str(i), + embedding=[rnd.uniform(-1, 1) for _ in range(16)], + embedder=embedder, + ) + for i in range(30) + ] + + for seed in range(5): + for lambda_mult in (0.0, 0.5, 1.0): + selected = [doc.id for doc in MMRReranker(lambda_mult=lambda_mult, top_n=8).rerank("q", build(seed))] + assert selected == reference(query_embedding, build(seed), lambda_mult, 8) + + +def test_mmr_scores_survive_the_search_api_schema(): + # /knowledge/search serializes results through VectorSearchResult, whose + # reranking_score bound must admit the negative scores MMR produces routinely. + schemas = pytest.importorskip("agno.os.routers.knowledge.schemas") + + results = MMRReranker(lambda_mult=0.5).rerank("q", _documents()) + + assert any(doc.reranking_score < 0 for doc in results) + for document in results: + schemas.VectorSearchResult.from_document(document) diff --git a/libs/agno/tests/unit/models/test_tool_result_falsy_values.py b/libs/agno/tests/unit/models/test_tool_result_falsy_values.py new file mode 100644 index 00000000000..37cef8b3b0c --- /dev/null +++ b/libs/agno/tests/unit/models/test_tool_result_falsy_values.py @@ -0,0 +1,87 @@ +"""Tool results that are falsy but meaningful (0, False, [], 0.0) must reach the +model as their string form on both the sync and the async execution paths.""" + +from typing import Any, AsyncIterator, Iterator, List + +import pytest + +from agno.models.base import Model +from agno.models.message import Message +from agno.models.response import ModelResponse +from agno.tools.function import Function, FunctionCall + + +class _StubModel(Model): + def __init__(self): + super().__init__(id="stub", name="stub", provider="stub") + + def invoke(self, *args, **kwargs) -> ModelResponse: + raise NotImplementedError + + async def ainvoke(self, *args, **kwargs) -> ModelResponse: + raise NotImplementedError + + def invoke_stream(self, *args, **kwargs) -> Iterator[ModelResponse]: + raise NotImplementedError + + async def ainvoke_stream(self, *args, **kwargs) -> AsyncIterator[ModelResponse]: + raise NotImplementedError + + def _parse_provider_response(self, response: Any, **kwargs) -> ModelResponse: + raise NotImplementedError + + def _parse_provider_response_delta(self, response: Any) -> ModelResponse: + raise NotImplementedError + + +def _function_call(return_value: Any) -> FunctionCall: + def tool() -> Any: + """Return a fixed value.""" + return return_value + + function = Function.from_callable(tool) + function.process_entrypoint() + return FunctionCall(function=function, arguments={}, call_id="call_1") + + +def _sync_tool_message(return_value: Any) -> Message: + results: List[Message] = [] + for _ in _StubModel().run_function_call(_function_call(return_value), function_call_results=results): + pass + assert len(results) == 1 + return results[0] + + +async def _async_tool_message(return_value: Any) -> Message: + results: List[Message] = [] + async for _ in _StubModel().arun_function_calls([_function_call(return_value)], function_call_results=results): + pass + assert len(results) == 1 + return results[0] + + +@pytest.mark.parametrize( + "return_value, expected", + [(0, "0"), (0.0, "0.0"), (False, "False"), ([], "[]"), ({}, "{}"), (None, "None"), (1, "1"), ("text", "text")], +) +def test_sync_tool_result_keeps_falsy_values(return_value, expected): + message = _sync_tool_message(return_value) + assert message.role == "tool" + assert message.content == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "return_value, expected", + [(0, "0"), (0.0, "0.0"), (False, "False"), ([], "[]"), ({}, "{}"), (1, "1"), ("text", "text")], +) +async def test_async_tool_result_keeps_falsy_values(return_value, expected): + message = await _async_tool_message(return_value) + assert message.role == "tool" + assert message.content == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("return_value", [0, False, [], None, "text"]) +async def test_sync_and_async_tool_results_match(return_value): + assert _sync_tool_message(return_value).content == (await _async_tool_message(return_value)).content diff --git a/libs/agno/tests/unit/models/yapi/__init__.py b/libs/agno/tests/unit/models/yapi/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/libs/agno/tests/unit/models/yapi/test_yapi.py b/libs/agno/tests/unit/models/yapi/test_yapi.py new file mode 100644 index 00000000000..504b2f1154e --- /dev/null +++ b/libs/agno/tests/unit/models/yapi/test_yapi.py @@ -0,0 +1,57 @@ +import os +from unittest.mock import patch + +import pytest + +from agno.exceptions import ModelAuthenticationError +from agno.models.utils import get_model, get_model_from_dict +from agno.models.yapi import YAPI + + +def test_yapi_initialization_with_api_key(): + model = YAPI(id="deepseek/deepseek-v4-flash", api_key="test-api-key") + assert model.id == "deepseek/deepseek-v4-flash" + assert model.api_key == "test-api-key" + assert model.base_url == "https://api.y-api.bestvirtualgoods.com/v1" + + +def test_yapi_initialization_without_api_key(): + with patch.dict(os.environ, {}, clear=True): + model = YAPI(id="deepseek/deepseek-v4-flash") + client_params = None + with pytest.raises(ModelAuthenticationError): + client_params = model._get_client_params() + assert client_params is None + + +def test_yapi_initialization_with_env_api_key(): + with patch.dict(os.environ, {"YAPI_API_KEY": "env-api-key"}): + model = YAPI(id="deepseek/deepseek-v4-flash") + assert model.api_key == "env-api-key" + + +def test_yapi_client_params(): + model = YAPI(id="deepseek/deepseek-v4-flash", api_key="test-api-key") + client_params = model._get_client_params() + assert client_params["api_key"] == "test-api-key" + assert client_params["base_url"] == "https://api.y-api.bestvirtualgoods.com/v1" + + +def test_yapi_default_values(): + model = YAPI(api_key="test-api-key") + assert model.id == "deepseek/deepseek-v4-flash" + assert model.name == "YAPI" + assert model.provider == "YAPI" + + +def test_yapi_resolves_from_string_syntax(): + model = get_model("yapi:deepseek/deepseek-v4-flash") + assert isinstance(model, YAPI) + assert model.id == "deepseek/deepseek-v4-flash" + + +def test_yapi_round_trips_through_dict(): + model = YAPI(id="z-ai/glm-5.3", api_key="test-api-key") + rebuilt = get_model_from_dict(model.to_dict()) + assert isinstance(rebuilt, YAPI) + assert rebuilt.id == "z-ai/glm-5.3" diff --git a/libs/agno/tests/unit/os/interfaces/test_agui_hitl.py b/libs/agno/tests/unit/os/interfaces/test_agui_hitl.py index 6b95f278b0d..1dab4c4e11b 100644 --- a/libs/agno/tests/unit/os/interfaces/test_agui_hitl.py +++ b/libs/agno/tests/unit/os/interfaces/test_agui_hitl.py @@ -9,6 +9,8 @@ """ import json +import logging +from typing import Any, AsyncIterator, Dict, Iterator, List from unittest.mock import MagicMock import pytest @@ -17,10 +19,16 @@ from ag_ui.core.types import Tool as AGUITool from ag_ui.core.types import ToolMessage as AGUIToolMessage +from fastapi.testclient import TestClient +from pydantic import ValidationError from agno.agent._tools import parse_tools from agno.agent.agent import Agent -from agno.models.response import ToolExecution, UserInputField +from agno.db.sqlite import SqliteDb +from agno.models.base import Model +from agno.models.response import ModelResponse, ModelResponseEvent, ToolExecution, UserInputField +from agno.os import AgentOS +from agno.os.interfaces.agui import AGUI from agno.os.interfaces.agui.handlers import on_run_completed from agno.os.interfaces.agui.input import parse_client_tools from agno.os.interfaces.agui.resume import ( @@ -34,6 +42,7 @@ ) from agno.run import RunContext from agno.run.agent import RunPausedEvent +from agno.run.base import RunStatus from agno.run.requirement import RunRequirement from agno.run.team import RunPausedEvent as TeamRunPausedEvent from agno.tools import tool @@ -52,6 +61,16 @@ def _tm(tool_call_id: str, content: str) -> AGUIToolMessage: return AGUIToolMessage(id="m-" + tool_call_id, role="tool", content=content, tool_call_id=tool_call_id) +def _tm_parts(tool_call_id: str, parts: List[Dict[str, Any]]) -> AGUIToolMessage: + """A tool message whose content is a list of content parts, the form ag-ui-protocol 1.0 added.""" + try: + return AGUIToolMessage.model_validate( + {"id": "m-" + tool_call_id, "role": "tool", "content": parts, "toolCallId": tool_call_id} + ) + except ValidationError: + pytest.skip("installed ag-ui-protocol only accepts string tool content") + + def _team_paused(*, requirements=None, tools=None) -> TeamRunPausedEvent: return TeamRunPausedEvent(requirements=requirements, tools=tools) @@ -224,6 +243,40 @@ def test_user_feedback_empty_selections_not_resolved(self): assert req.is_resolved() is False +class TestPauseResolutionWithContentParts: + """ag-ui-protocol 1.0 lets a tool message carry a list of content parts instead of a string. + The answer is the text of its text parts, whichever pause type it resolves.""" + + def test_confirmation_in_a_text_part_confirms(self): + req = RunRequirement(ToolExecution(tool_call_id="tc1", tool_name="x", requires_confirmation=True)) + answer = _tm_parts("tc1", [{"type": "text", "text": json.dumps({"accepted": True})}]) + resolve_requirements_from_tool_messages([req], [answer]) + assert req.tool_execution.confirmed is True + assert req.is_resolved() + + def test_external_execution_result_is_the_text_of_the_parts(self): + te = ToolExecution(tool_call_id="tc3", tool_name="run", external_execution_required=True) + req = RunRequirement(te) + answer = _tm_parts("tc3", [{"type": "text", "text": "first line"}, {"type": "text", "text": "second line"}]) + resolve_requirements_from_tool_messages([req], [answer]) + assert req.external_execution_result == "first line\nsecond line" + + def test_external_execution_drops_media_parts_with_a_warning(self, caplog): + te = ToolExecution(tool_call_id="tc4", tool_name="run", external_execution_required=True) + req = RunRequirement(te) + answer = _tm_parts( + "tc4", + [ + {"type": "text", "text": "the chart"}, + {"type": "image", "source": {"type": "url", "value": "https://example.com/chart.png"}}, + ], + ) + with caplog.at_level(logging.WARNING, logger="agno"): + resolve_requirements_from_tool_messages([req], [answer]) + assert req.external_execution_result == "the chart" + assert any("tc4" in record.message and "image" in record.message for record in caplog.records) + + class TestDedupe: def test_backend_confirmation_tool_wins_over_same_named_client_tool(self): """A frontend-advertised client tool must NOT shadow the agent's own @@ -362,3 +415,83 @@ def test_team_member_requirement_resolves_via_existing_merge(self): resolve_requirements_from_tool_messages([req], [_tm("m-res", json.dumps({"accepted": True}))]) assert req.tool_execution.confirmed is True assert req.is_resolved() + + +def _sse_events(text: str) -> List[Dict[str, Any]]: + return [json.loads(line[5:]) for line in text.splitlines() if line.startswith("data:")] + + +class _ScriptedModel(Model): + """Calls change_background on its first turn and answers in text afterwards. Records the tool results it is sent.""" + + def __init__(self): + super().__init__(id="scripted", name="scripted", provider="test") + self.turns = 0 + self.tool_results_seen: List[Any] = [] + + def _next(self, messages: List[Any]) -> ModelResponse: + self.turns += 1 + self.tool_results_seen.extend(m.content for m in messages if m.role == "tool") + if self.turns == 1: + function = {"name": "change_background", "arguments": json.dumps({"color": "blue"})} + return ModelResponse( + role="assistant", tool_calls=[{"id": "call_1", "type": "function", "function": function}] + ) + return ModelResponse(role="assistant", content="all done", event=ModelResponseEvent.assistant_response.value) + + def invoke(self, messages=None, *args, **kwargs) -> ModelResponse: + return self._next(messages or []) + + async def ainvoke(self, messages=None, *args, **kwargs) -> ModelResponse: + return self._next(messages or []) + + def invoke_stream(self, messages=None, *args, **kwargs) -> Iterator[ModelResponse]: + yield self._next(messages or []) + + async def ainvoke_stream(self, messages=None, *args, **kwargs) -> AsyncIterator[ModelResponse]: + yield self._next(messages or []) + + def _parse_provider_response(self, response: Any, **kwargs) -> ModelResponse: + return response + + def _parse_provider_response_delta(self, response: Any) -> ModelResponse: + return response + + +class TestResumeWithContentPartsThroughTheRoute: + """A pause answered over POST /agui with list-form tool content must finish AND be saved as finished. + SqliteDb, not InMemoryDb: the defect was a run that could not be serialized on save, and only a + database that serializes the run can show it.""" + + def test_external_execution_answer_reaches_the_model_as_text_and_the_run_is_saved_completed(self, tmp_path): + answer = [{"type": "text", "text": "blue is set"}] + _tm_parts("probe", answer) # skips on an ag-ui-protocol that only accepts string tool content + model = _ScriptedModel() + agent = Agent(id="parts-agent", model=model, db=SqliteDb(db_file=str(tmp_path / "parts.db")), telemetry=False) + client = TestClient(AgentOS(agents=[agent], interfaces=[AGUI(agent=agent)], telemetry=False).get_app()) + frontend_tool = { + "name": "change_background", + "description": "Change the page background", + "parameters": {"type": "object", "properties": {"color": {"type": "string"}}}, + } + + def post(run_id: str, messages: list): + body = {"threadId": "thread-parts", "runId": run_id, "state": {}, "messages": messages} + return client.post("/agui", json={**body, "tools": [frontend_tool], "context": [], "forwardedProps": {}}) + + user = {"id": "u1", "role": "user", "content": "go"} + paused = post("run-1", [user]) + call = next(e for e in _sse_events(paused.text) if e["type"] == "TOOL_CALL_START") + function = {"name": "change_background", "arguments": json.dumps({"color": "blue"})} + assistant = { + "id": "a1", + "role": "assistant", + "toolCalls": [{"id": call["toolCallId"], "type": "function", "function": function}], + } + tool_message = {"id": "t1", "role": "tool", "toolCallId": call["toolCallId"], "content": answer} + resumed = post("run-2", [user, assistant, tool_message]) + + assert resumed.status_code == 200 + assert [e["type"] for e in _sse_events(resumed.text)][-1] == "RUN_FINISHED" + assert model.tool_results_seen == ["blue is set"] + assert [run.status for run in agent.get_session(session_id="thread-parts").runs] == [RunStatus.completed] diff --git a/libs/agno/tests/unit/os/test_validation_error_body.py b/libs/agno/tests/unit/os/test_validation_error_body.py index 4186cc959c0..abc3fe32a46 100644 --- a/libs/agno/tests/unit/os/test_validation_error_body.py +++ b/libs/agno/tests/unit/os/test_validation_error_body.py @@ -116,8 +116,9 @@ def test_agui_dependency_validator_is_422_with_message(self, tmp_path): agent = Agent(id="qa-agent", name="QA Agent", db=db) agent_os = AgentOS(agents=[agent], db=db, telemetry=False, interfaces=[AGUI(agent=agent)]) client = TestClient(agent_os.get_app(), raise_server_exceptions=False) - # A binary content item with no id, url, or data trips BinaryInputContent's - # model_validator inside the ag_ui dependency. + # A binary content item with no id, url, or data. ag-ui-protocol before 1.0 rejects it + # from BinaryInputContent's model_validator (a ValueError); 1.0 dropped the binary part + # and rejects it as an unknown union tag, so only pre-1.0 SDKs reach the ValueError path here. bad_content = [{"type": "binary", "mimeType": "application/octet-stream"}] resp = client.post( "/agui", @@ -132,9 +133,15 @@ def test_agui_dependency_validator_is_422_with_message(self, tmp_path): }, ) assert resp.status_code == 422, f"expected 422, got {resp.status_code}: {resp.text[:200]}" - assert "BinaryInputContent requires id, url, or data" in resp.text, ( - f"message missing from body: {resp.text[:300]}" + detail = resp.json()["detail"] + assert detail and all(err["loc"][:3] == ["body", "messages", 0] for err in detail), ( + f"errors do not name the offending message: {resp.text[:300]}" ) + for err in detail: + if err["type"] == "value_error": + assert "BinaryInputContent requires id, url, or data" in err["msg"], ( + f"validator message missing from body: {resp.text[:300]}" + ) class TestOwnedAndBorrowedAppsAgree: diff --git a/libs/agno/tests/unit/os/test_ws_workflow_submission_parity.py b/libs/agno/tests/unit/os/test_ws_workflow_submission_parity.py new file mode 100644 index 00000000000..d1508cab14d --- /dev/null +++ b/libs/agno/tests/unit/os/test_ws_workflow_submission_parity.py @@ -0,0 +1,169 @@ +"""Workflow WebSocket submissions apply the same admission rules as HTTP. + +The HTTP and WebSocket doors share the durable core (validation, enqueue, +register, prepare, tail), but the WebSocket door drifted on the guards that +run before it: which submissions may ride the queue, and which sessions a +caller may write into. +""" + +import json +from types import SimpleNamespace +from typing import Any, List + +import pytest + +from agno.db.schemas.scheduler import COMPONENT_VERSION_METADATA_KEY + + +class FakeWebSocket: + def __init__(self, app_state: Any): + self.sent: List[dict] = [] + self.app = SimpleNamespace(state=app_state) + + async def send_text(self, text: str) -> None: + self.sent.append(json.loads(text)) + + +@pytest.fixture +def ws_env(monkeypatch): + from agno.db.in_memory import InMemoryDb + from agno.job_queue.config import QueueConfig + from agno.job_queue.store import InMemoryQueueStore + from agno.os.event_streams.in_memory import InMemoryEventStream + from agno.os.managers import EventsBuffer, SSESubscriberManager + from agno.os.routers.workflows import router as ws_router + from agno.workflow.workflow import Workflow + + stream = InMemoryEventStream(events_buffer=EventsBuffer(), subscriber_manager=SSESubscriberManager()) + monkeypatch.setattr(ws_router, "get_event_stream", lambda: stream) + workflow = Workflow(id="wf1", name="WF", db=InMemoryDb()) + monkeypatch.setattr(ws_router, "get_workflow_by_id", lambda **kwargs: workflow) + + async def no_prepare(*args, **kwargs): + return None + + monkeypatch.setattr(ws_router, "aprepare_accepted_or_abort", no_prepare) + # The non-durable path hands the run to the workflow itself; record it + arun_calls: List[dict] = [] + + async def recording_arun(**kwargs): + arun_calls.append(kwargs) + return None + + monkeypatch.setattr(workflow, "arun", recording_arun) + store = InMemoryQueueStore() + ws = FakeWebSocket(SimpleNamespace(queue_worker=SimpleNamespace(store=store, config=QueueConfig(durable=True)))) + os_stub = SimpleNamespace(workflows=[workflow], db=None, registry=None) + yield SimpleNamespace( + router=ws_router, stream=stream, ws=ws, os=os_stub, store=store, workflow=workflow, arun_calls=arun_calls + ) + + +def _queued_acks(env) -> List[dict]: + return [f for f in env.ws.sent if f.get("event") == "queued"] + + +@pytest.mark.asyncio +async def test_version_pinned_submission_does_not_ride_the_queue(ws_env): + """The worker resolves the registry instance, so a ticket cannot carry a + version pin. HTTP refuses to queue pinned submissions and runs them + in-process with the pin stamped on the run; the WebSocket door must do + the same instead of queueing the run and silently executing whatever + version is current.""" + from agno.run.base import RunStatus + + env = ws_env + await env.router.handle_workflow_via_websocket( + env.ws, {"workflow_id": "wf1", "session_id": "s1", "message": "hi", "version": 2}, env.os + ) + try: + assert not _queued_acks(env), "a version-pinned submission must not be queued" + assert await env.store.count_queued_jobs() == 0 + assert len(env.arun_calls) == 1, "it must take the in-process path instead" + stamped = env.arun_calls[0].get("metadata") or {} + assert stamped.get(COMPONENT_VERSION_METADATA_KEY) == 2, "and carry the pin on the run" + finally: + for ack in _queued_acks(env): + await env.stream.complete_run(ack["run_id"], RunStatus.completed) + await env.router.cancel_subscription_pump(env.ws) + + +def _isolated(ws_router): + return ws_router.WebSocketAuthContext(jwt_enabled=True, is_admin=False, user_isolation_enabled=True) + + +@pytest.mark.asyncio +async def test_submission_into_another_users_session_is_refused(ws_env, monkeypatch): + """HTTP refuses a run into a session owned by someone else before any + background work: the runs table has no ownership predicate, so an + unguarded write lands in the owner's history as their own turn. The + WebSocket door pins the caller's identity to the token but let the + client choose any session id; it must apply the same guard.""" + from agno.run.base import RunStatus + + env = ws_env + monkeypatch.setattr(env.workflow.db, "get_session", lambda **kwargs: {"session_id": "s1", "user_id": "owner"}) + await env.router.handle_workflow_via_websocket( + env.ws, + {"workflow_id": "wf1", "session_id": "s1", "message": "hi"}, + env.os, + ws_user_context={"user_id": "intruder"}, + ws_auth=_isolated(env.router), + ) + try: + assert not _queued_acks(env), "a run into another user's session must not be queued" + assert await env.store.count_queued_jobs() == 0 + assert not env.arun_calls, "nor executed in-process" + errors = [f for f in env.ws.sent if f.get("event") == "error"] + assert errors, "the caller must be told the submission was refused" + finally: + for ack in _queued_acks(env): + await env.stream.complete_run(ack["run_id"], RunStatus.completed) + await env.router.cancel_subscription_pump(env.ws) + + +@pytest.mark.asyncio +async def test_owner_submission_into_own_session_is_queued(ws_env, monkeypatch): + from agno.run.base import RunStatus + + env = ws_env + monkeypatch.setattr(env.workflow.db, "get_session", lambda **kwargs: {"session_id": "s1", "user_id": "owner"}) + await env.router.handle_workflow_via_websocket( + env.ws, + {"workflow_id": "wf1", "session_id": "s1", "message": "hi"}, + env.os, + ws_user_context={"user_id": "owner"}, + ws_auth=_isolated(env.router), + ) + try: + acks = _queued_acks(env) + assert len(acks) == 1, f"the owner's submission must be accepted, got {env.ws.sent}" + finally: + for ack in _queued_acks(env): + await env.stream.complete_run(ack["run_id"], RunStatus.completed) + await env.router.cancel_subscription_pump(env.ws) + + +@pytest.mark.asyncio +async def test_submission_without_session_id_gets_a_fresh_session(ws_env): + """HTTP mints a new session for a submission that names none. The + WebSocket door fell back to the workflow's own session_id first, so every + client omitting the field on a workflow configured with one pooled into + a single session, and under per-session queueing they would all line up + behind each other.""" + from agno.run.base import RunStatus + + env = ws_env + env.workflow.session_id = "fixed-on-the-workflow" + await env.router.handle_workflow_via_websocket(env.ws, {"workflow_id": "wf1", "message": "one"}, env.os) + await env.router.handle_workflow_via_websocket(env.ws, {"workflow_id": "wf1", "message": "two"}, env.os) + try: + acks = _queued_acks(env) + assert len(acks) == 2 + sessions = {ack["session_id"] for ack in acks} + assert "fixed-on-the-workflow" not in sessions, "the workflow's own session_id is not a default for clients" + assert len(sessions) == 2, "each submission without a session_id gets its own session, as over HTTP" + finally: + for ack in _queued_acks(env): + await env.stream.complete_run(ack["run_id"], RunStatus.completed) + await env.router.cancel_subscription_pump(env.ws) diff --git a/libs/agno/tests/unit/reader/test_csv_reader.py b/libs/agno/tests/unit/reader/test_csv_reader.py index f6b77d0b765..00181305ce2 100644 --- a/libs/agno/tests/unit/reader/test_csv_reader.py +++ b/libs/agno/tests/unit/reader/test_csv_reader.py @@ -4,6 +4,7 @@ import pytest +from agno.knowledge.chunking.row import RowChunking from agno.knowledge.document.base import Document from agno.knowledge.reader.csv_reader import CSVReader @@ -227,6 +228,21 @@ async def test_async_read_multi_page_csv(csv_reader, multi_page_csv_file): assert documents[10].meta_data["rows"] == 1 +@pytest.mark.asyncio +@pytest.mark.parametrize("skip_header", [False, True]) +@pytest.mark.parametrize("page_size", [1, 5, 1000]) +async def test_async_read_multi_page_csv_preserves_rows_and_numbers(multi_page_csv_file, skip_header, page_size): + reader = CSVReader(chunking_strategy=RowChunking(skip_header=skip_header)) + + sync_documents = reader.read(multi_page_csv_file) + async_documents = await reader.async_read(multi_page_csv_file, page_size=page_size) + + assert [document.content for document in async_documents] == [document.content for document in sync_documents] + assert [document.meta_data["row_number"] for document in async_documents] == list( + range(2 if skip_header else 1, 12) + ) + + @pytest.mark.asyncio async def test_async_read_with_chunking(csv_reader, csv_file): async def mock_achunk(doc): diff --git a/libs/agno/tests/unit/reader/test_firecrawl_reader.py b/libs/agno/tests/unit/reader/test_firecrawl_reader.py index 5aced00c70b..1961dd342b2 100644 --- a/libs/agno/tests/unit/reader/test_firecrawl_reader.py +++ b/libs/agno/tests/unit/reader/test_firecrawl_reader.py @@ -52,8 +52,8 @@ def test_scrape_basic(mock_scrape_response): assert len(documents) == 1 assert documents[0].name == "https://example.com" assert documents[0].id == "https://example.com_1" - # Content is joined with spaces instead of newlines - expected_content = "# Test Website This is test content from a scraped website." + # Repeated newlines collapse to a single newline + expected_content = "# Test Website\nThis is test content from a scraped website." assert documents[0].content == expected_content # Verify FirecrawlApp was called correctly @@ -186,9 +186,9 @@ def test_crawl_basic(mock_crawl_response): assert len(documents) == 2 # Base URL is used for name assert documents[0].name == "https://example.com" - # Content joined with spaces - assert documents[0].content == "# Page 1 This is content from page 1." - assert documents[1].content == "# Page 2 This is content from page 2." + # Repeated newlines collapse to a single newline + assert documents[0].content == "# Page 1\nThis is content from page 1." + assert documents[1].content == "# Page 2\nThis is content from page 2." # Verify FirecrawlApp was called correctly MockFirecrawlApp.assert_called_once_with(api_key=None) @@ -261,7 +261,7 @@ def test_read_scrape_mode(mock_scrape_response): documents = reader.read("https://example.com") assert len(documents) == 1 - expected_content = "# Test Website This is test content from a scraped website." + expected_content = "# Test Website\nThis is test content from a scraped website." assert documents[0].content == expected_content mock_app.scrape_url.assert_called_once() diff --git a/libs/agno/tests/unit/reader/test_sitemap_reader.py b/libs/agno/tests/unit/reader/test_sitemap_reader.py index c67ce30e159..7cc5d81aa78 100644 --- a/libs/agno/tests/unit/reader/test_sitemap_reader.py +++ b/libs/agno/tests/unit/reader/test_sitemap_reader.py @@ -285,6 +285,40 @@ def test_gzipped_sitemap_parsed(): assert documents[0].meta_data["source"] == "sitemap" +@pytest.mark.parametrize( + "invalid_gzip", + [ + b"\x1f\x8b", + gzip.compress(b"", mtime=0)[:-1], + # A gzip header followed by a reserved DEFLATE block type + b"\x1f\x8b\x08\x00\x00\x00\x00\x00\x00\xff\x07", + ], + ids=["truncated-header", "truncated-trailer", "invalid-deflate"], +) +@pytest.mark.parametrize("asynchronous", [False, True]) +def test_invalid_gzip_sitemap_tries_next_candidate(invalid_gzip, asynchronous): + # A gzipped body that does not decompress is not a sitemap; the next candidate is tried + routes = { + "https://example.com/sitemap.xml.gz": (invalid_gzip, "application/gzip"), + "https://example.com/sitemap.xml": (urlset_xml("https://example.com/page-a"), "application/xml"), + "https://example.com/page-a": (html_page("A", "Alpha content"), "text/html"), + } + reader = make_reader() + with mock_site(routes) as requested: + if asynchronous: + documents = asyncio.run(reader.async_read("https://example.com/sitemap.xml.gz")) + else: + documents = reader.read("https://example.com/sitemap.xml.gz") + + assert "https://example.com/sitemap.xml.gz" in requested + assert "https://example.com/sitemap.xml" in requested + assert len(documents) == 1 + assert documents[0].content == "Alpha content" + assert documents[0].meta_data["url"] == "https://example.com/page-a" + assert documents[0].meta_data["source"] == "sitemap" + assert "discovery_incomplete" not in documents[0].meta_data + + def test_nested_index_cycle_terminates(): routes = { "https://example.com/sitemap.xml": (sitemapindex_xml("https://example.com/idx2.xml"), "application/xml"), @@ -535,8 +569,8 @@ def test_source_header_lands_in_first_chunk_only(): documents = reader.read("https://example.com/sitemap.xml") assert len(documents) >= 2 - # FixedSizeChunking collapses whitespace, so the header's newlines become spaces - assert documents[0].content.startswith("# Page A Source: https://example.com/a ") + # FixedSizeChunking collapses repeated newlines, so the header's blank line becomes a single newline + assert documents[0].content.startswith("# Page A\nSource: https://example.com/a\n") for later in documents[1:]: assert "# Page A" not in later.content assert "Source:" not in later.content @@ -628,6 +662,41 @@ def test_failed_index_child_marks_documents_discovery_incomplete(): ) +@pytest.mark.parametrize( + "invalid_gzip", + [ + b"\x1f\x8b", + gzip.compress(b"", mtime=0)[:-1], + # A gzip header followed by a reserved DEFLATE block type + b"\x1f\x8b\x08\x00\x00\x00\x00\x00\x00\xff\x07", + ], + ids=["truncated-header", "truncated-trailer", "invalid-deflate"], +) +@pytest.mark.parametrize("asynchronous", [False, True]) +def test_invalid_gzip_index_child_preserves_healthy_pages(invalid_gzip, asynchronous): + routes = { + "https://example.com/sitemap.xml": ( + sitemapindex_xml("https://example.com/sitemap-a.xml", "https://example.com/sitemap-b.xml.gz"), + "application/xml", + ), + "https://example.com/sitemap-a.xml": (urlset_xml("https://example.com/page-a"), "application/xml"), + # sitemap-b.xml.gz does not decompress: an entire shard is missing from this read + "https://example.com/sitemap-b.xml.gz": (invalid_gzip, "application/gzip"), + "https://example.com/page-a": (html_page("A", "Alpha content"), "text/html"), + } + reader = make_reader() + with mock_site(routes): + if asynchronous: + documents = asyncio.run(reader.async_read("https://example.com/sitemap.xml")) + else: + documents = reader.read("https://example.com/sitemap.xml") + + assert len(documents) == 1 + assert documents[0].content == "Alpha content" + assert documents[0].meta_data["url"] == "https://example.com/page-a" + assert documents[0].meta_data["discovery_incomplete"] is True + + def test_complete_discovery_carries_no_incomplete_flag(): routes = { "https://example.com/sitemap.xml": (urlset_xml("https://example.com/page-a"), "application/xml"), diff --git a/libs/agno/tests/unit/reader/test_tavily_reader.py b/libs/agno/tests/unit/reader/test_tavily_reader.py index 05cbb1bbeed..2f3546ecfd7 100644 --- a/libs/agno/tests/unit/reader/test_tavily_reader.py +++ b/libs/agno/tests/unit/reader/test_tavily_reader.py @@ -53,8 +53,8 @@ def test_extract_basic(mock_extract_response): assert len(documents) == 1 assert documents[0].name == "https://example.com" assert documents[0].id == "https://example.com_1" - # Content is joined with spaces instead of newlines - expected_content = "# Test Website This is test content from an extracted website." + # Repeated newlines collapse to a single newline + expected_content = "# Test Website\nThis is test content from an extracted website." assert documents[0].content == expected_content # Verify TavilyClient was called correctly @@ -242,7 +242,7 @@ def test_read_method(mock_extract_response): documents = reader.read("https://example.com") assert len(documents) == 1 - expected_content = "# Test Website This is test content from an extracted website." + expected_content = "# Test Website\nThis is test content from an extracted website." assert documents[0].content == expected_content mock_client.extract.assert_called_once() diff --git a/libs/agno/tests/unit/reader/test_text_reader.py b/libs/agno/tests/unit/reader/test_text_reader.py index 7af0a34ba81..93e4c263e68 100644 --- a/libs/agno/tests/unit/reader/test_text_reader.py +++ b/libs/agno/tests/unit/reader/test_text_reader.py @@ -1,5 +1,6 @@ import asyncio -from io import BytesIO +from contextlib import ExitStack +from io import BytesIO, StringIO from pathlib import Path from typing import List from unittest.mock import patch @@ -39,6 +40,35 @@ def test_read_text_bytesio(): assert documents[0].content == test_data +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True], ids=["sync", "async"]) +@pytest.mark.parametrize("chunk", [False, True], ids=["whole", "chunked"]) +@pytest.mark.parametrize("stream_type", ["stringio", "text_file", "utf8_bytes", "latin1_bytes"]) +async def test_read_text_and_binary_streams(tmp_path, use_async, chunk, stream_type): + content = "Agent notes: café." if stream_type == "latin1_bytes" else "Agent notes: café 中文." + encoding = None if stream_type == "utf8_bytes" else "latin-1" + + with ExitStack() as stack: + if stream_type == "text_file": + path = tmp_path / "notes.txt" + path.write_text(content, encoding="utf-8") + stream = stack.enter_context(path.open(encoding="utf-8")) + elif stream_type == "stringio": + stream = stack.enter_context(StringIO(content)) + else: + stream = stack.enter_context(BytesIO(content.encode(encoding or "utf-8"))) + + # Reading must rewind the input; text streams are already decoded. + stream.read(3) + reader = TextReader(chunk=chunk, encoding=encoding) + documents = await reader.async_read(stream, name="notes") if use_async else reader.read(stream, name="notes") + + assert len(documents) == 1 + assert documents[0].content == content + assert documents[0].name == "notes" + assert not stream.closed + + def test_chunking(): # Test document chunking functionality test_data = "Hello, world!" diff --git a/libs/agno/tests/unit/reasoning/test_reasoning_checkers.py b/libs/agno/tests/unit/reasoning/test_reasoning_checkers.py index bdde2731e20..a2bc13b0fbb 100644 --- a/libs/agno/tests/unit/reasoning/test_reasoning_checkers.py +++ b/libs/agno/tests/unit/reasoning/test_reasoning_checkers.py @@ -247,6 +247,43 @@ def test_openai_like_with_deepseek_r1(): assert is_openai_reasoning_model(model) is True +@pytest.mark.parametrize( + "model_id", + ["deepseek-v4-pro", "deepseek-v4-flash", "deepseek-v3.1-terminus", "deepseek-v3.2-exp", "deepseek-reasoner"], +) +def test_openai_like_with_deepseek_thinking_ids(model_id): + """DeepSeek thinking-mode IDs served over OpenAI-compatible endpoints are native reasoning models (#10277).""" + from agno.models.openai.like import OpenAILike + + model = OpenAILike( + id=model_id, + name="DeepSeek", + ) + assert is_openai_reasoning_model(model) is True + + +def test_dashscope_with_deepseek_v4(): + """Test DashScope (an OpenAILike subclass) with a DeepSeek V4 reasoning_model ID returns True.""" + from agno.models.dashscope.dashscope import DashScope + + model = DashScope( + id="deepseek-v4-pro", + name="DashScope", + ) + assert is_openai_reasoning_model(model) is True + + +def test_openai_like_with_deepseek_chat_stays_false(): + """Test OpenAILike model with the non-reasoning deepseek-chat ID stays False.""" + from agno.models.openai.like import OpenAILike + + model = OpenAILike( + id="deepseek-chat", + name="DeepSeek", + ) + assert is_openai_reasoning_model(model) is False + + def test_openai_chat_without_reasoning_id(): """Test OpenAIChat model without reasoning model ID returns False.""" model = MockModel( diff --git a/libs/agno/tests/unit/tools/models/test_aimlapi.py b/libs/agno/tests/unit/tools/models/test_aimlapi.py new file mode 100644 index 00000000000..af796a0cb68 --- /dev/null +++ b/libs/agno/tests/unit/tools/models/test_aimlapi.py @@ -0,0 +1,384 @@ +import json +import subprocess +import sys +from typing import Any, Dict, List +from unittest.mock import patch + +import httpx +import pytest + +from agno.models.aimlapi.constants import AIMLAPI_HEADERS +from agno.tools.function import ToolResult +from agno.tools.models.aimlapi import AIMLAPITools + + +class Gateway: + """Records every request and answers with the gateway's documented shapes.""" + + def __init__(self): + self.calls: List[httpx.Request] = [] + self.video_statuses = ["queued", "generating", "completed"] + self.stt_statuses = ["queued", "completed"] + self.poll_failures: List[int] = [] # HTTP statuses to answer polls with, before the real one + self.asset_content_type = "image/png" + self.transcript: Any = "hello from agno" + self.stt_error: Any = {"name": "ProviderError", "message": "transcription failed"} + self.error_shape: Any = {"message": "content policy"} + + def handle(self, request: httpx.Request) -> httpx.Response: + self.calls.append(request) + host, path = request.url.host, request.url.path + if host == "cdn.example": + assert "authorization" not in request.headers, "asset downloads must not carry the account key" + if path.endswith(".mp4"): + return httpx.Response(200, content=b"\x00mp4", headers={"content-type": "video/mp4"}) + if path.endswith(".mp3"): + return httpx.Response(200, content=b"\x00mp3", headers={"content-type": "audio/mpeg"}) + return httpx.Response(200, content=b"\x89PNG", headers={"content-type": self.asset_content_type}) + if path == "/v1/images/generations": + return httpx.Response(200, json={"data": [{"url": "https://cdn.example/out.png"}]}) + if path == "/v2/video/generations" and request.method == "POST": + return httpx.Response(200, json={"id": "gen-1", "status": self.video_statuses.pop(0)}) + if path == "/v2/video/generations": + if self.poll_failures: + return httpx.Response(self.poll_failures.pop(0), json={"message": "try later"}) + status = self.video_statuses.pop(0) + body: Dict[str, Any] = {"id": "gen-1", "status": status} + if status == "completed": + body["video"] = {"url": "https://cdn.example/out.mp4"} + if status == "error": + body["error"] = self.error_shape + return httpx.Response(200, json=body) + if path == "/v1/tts": + return httpx.Response(200, json={"audio": {"url": "https://cdn.example/out.mp3"}}) + if path == "/v1/stt/create": + return httpx.Response(200, json={"generation_id": "stt-1", "status": self.stt_statuses.pop(0)}) + if path == "/v1/stt/stt-1": + status = self.stt_statuses.pop(0) + body = {"generation_id": "stt-1", "status": status} + if status in ("error", "failed"): + body["error"] = self.stt_error + if status == "completed": + body["result"] = {"results": {"channels": [{"alternatives": [{"transcript": self.transcript}]}]}} + return httpx.Response(200, json=body) + return httpx.Response(404, json={"message": f"no route for {request.method} {path}"}) + + +@pytest.fixture +def gateway(): + gw = Gateway() + transport = httpx.MockTransport(lambda request: gw.handle(request)) + + def post(url, **kwargs): + with httpx.Client(transport=transport) as client: + return client.post(url, **kwargs) + + def get(url, **kwargs): + with httpx.Client(transport=transport) as client: + return client.get(url, **kwargs) + + real_async_client = httpx.AsyncClient + + def async_client(**kwargs): + kwargs.pop("timeout", None) + return real_async_client(transport=transport, **kwargs) + + with ( + patch("agno.tools.models.aimlapi.httpx.post", side_effect=post), + patch("agno.tools.models.aimlapi.httpx.get", side_effect=get), + patch("agno.tools.models.aimlapi.httpx.AsyncClient", side_effect=async_client), + patch("agno.tools.models.aimlapi.time.sleep"), + ): + yield gw + + +def tools(**kwargs) -> AIMLAPITools: + return AIMLAPITools(api_key="sk-test", **kwargs) + + +def paths(gateway: Gateway): + return [(c.method, c.url.path) for c in gateway.calls] + + +# --- construction -------------------------------------------------------------- + + +def test_reads_key_from_env(monkeypatch): + monkeypatch.setenv("AIMLAPI_API_KEY", "sk-env") + assert AIMLAPITools().api_key == "sk-env" + + +def test_requires_a_key(monkeypatch): + monkeypatch.delenv("AIMLAPI_API_KEY", raising=False) + with pytest.raises(ValueError, match="AIMLAPI_API_KEY not set"): + AIMLAPITools() + + +def test_rejects_an_unknown_speech_format(): + with pytest.raises(ValueError, match="speech_format"): + tools(speech_format="ogg") + + +def test_registers_every_tool_with_async_variants(): + t = tools() + assert set(t.functions) == {"generate_image", "generate_video", "generate_speech", "transcribe_audio"} + assert set(t.async_functions) == set(t.functions) + + +def test_flags_select_tools(): + t = tools(enable_generate_video=False, enable_generate_speech=False, enable_transcribe_audio=False) + assert list(t.functions) == ["generate_image"] + assert set(tools(enable_generate_image=False, all=True).functions) == { + "generate_image", + "generate_video", + "generate_speech", + "transcribe_audio", + } + + +def test_timeout_reaches_the_toolkit_and_the_requests(): + t = tools(timeout=15) + assert t.timeout == 15 + assert t.request_timeout == 15 + + +def test_accepts_the_chat_models_versioned_base_url(): + assert tools(base_url="https://api.aimlapi.com/v1").base_url == "https://api.aimlapi.com" + assert tools(base_url="https://api.aimlapi.com/v1/").base_url == "https://api.aimlapi.com" + assert tools(base_url="https://proxy.example/aimlapi").base_url == "https://proxy.example/aimlapi" + + +def test_importing_the_toolkit_does_not_need_openai(): + code = "import sys; sys.modules['openai'] = None; import agno.tools.models.aimlapi; print('ok')" + result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True) + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "ok" + + +# --- attribution --------------------------------------------------------------- + + +def test_attribution_headers_ride_calls_to_the_gateway(gateway): + tools().generate_image("a cat") + submit = gateway.calls[0] + assert submit.headers["authorization"] == "Bearer sk-test" + for key, value in AIMLAPI_HEADERS.items(): + assert submit.headers[key] == value + + +def test_attribution_headers_stay_off_a_proxy(gateway): + tools(base_url="https://proxy.example/aimlapi/").generate_speech("hi") + submit = gateway.calls[0] + assert submit.url.path == "/aimlapi/v1/tts" + assert submit.headers["authorization"] == "Bearer sk-test" + assert not any(key.lower().startswith("x-aimlapi-") for key in submit.headers) + + +# --- generate_image ------------------------------------------------------------ + + +def test_generate_image_downloads_the_asset(gateway): + result = tools(image_size="1024x1024").generate_image("a cat") + assert isinstance(result, ToolResult) + assert result.content == "Image generated successfully." + image = result.images[0] + assert (image.content, image.mime_type, image.format, image.original_prompt) == ( + b"\x89PNG", + "image/png", + "png", + "a cat", + ) + assert json.loads(gateway.calls[0].content) == { + "model": "openai/gpt-image-2", + "prompt": "a cat", + "size": "1024x1024", + } + + +@pytest.mark.asyncio +async def test_agenerate_image_matches_the_sync_tool(gateway): + result = await tools().agenerate_image("a cat") + assert result.content == "Image generated successfully." + assert result.images[0].format == "png" + assert paths(gateway) == [("POST", "/v1/images/generations"), ("GET", "/out.png")] + + +def test_octet_stream_assets_are_typed_from_the_url(gateway): + gateway.asset_content_type = "application/octet-stream" + image = tools().generate_image("a cat").images[0] + assert (image.mime_type, image.format) == ("image/png", "png") + + +def test_generate_image_reports_gateway_errors(gateway): + gateway.handle = lambda request: httpx.Response(400, json={"message": "Validation failed"}) + result = tools().generate_image("a cat") + assert result.content == "Failed to generate image: AI/ML API returned HTTP 400: Validation failed" + assert not result.images + + +def test_string_shaped_errors_are_readable(gateway): + gateway.handle = lambda request: httpx.Response(400, json={"error": "bad prompt"}) + assert tools().generate_image("a cat").content.endswith("HTTP 400: bad prompt") + + +# --- generate_video ------------------------------------------------------------ + + +def test_generate_video_submits_polls_and_collects(gateway): + result = tools(video_duration=4, video_resolution="480p").generate_video("a boat") + assert result.content == "Video generated successfully." + video = result.videos[0] + assert (video.content, video.mime_type, video.format) == (b"\x00mp4", "video/mp4", "mp4") + assert paths(gateway) == [ + ("POST", "/v2/video/generations"), + ("GET", "/v2/video/generations"), + ("GET", "/v2/video/generations"), + ("GET", "/out.mp4"), + ] + assert dict(gateway.calls[1].url.params) == {"generation_id": "gen-1"} + assert json.loads(gateway.calls[0].content) == { + "model": "bytedance/seedance-2-5", + "prompt": "a boat", + "duration": 4, + "resolution": "480p", + } + + +@pytest.mark.asyncio +async def test_agenerate_video_polls_without_blocking(gateway): + with patch("agno.tools.models.aimlapi.asyncio.sleep") as sleep: + result = await tools().agenerate_video("a boat") + assert result.content == "Video generated successfully." + assert result.videos[0].format == "mp4" + assert sleep.await_count == 2 + assert paths(gateway)[-1] == ("GET", "/out.mp4") + + +def test_generate_video_reports_a_failed_job(gateway): + gateway.video_statuses = ["queued", "error"] + result = tools().generate_video("a boat") + assert result.content == "Failed to generate video: content policy" + assert not result.videos + + +def test_generate_video_reports_a_string_error(gateway): + gateway.video_statuses = ["queued", "error"] + gateway.error_shape = "quota exhausted" + assert tools().generate_video("a boat").content == "Failed to generate video: quota exhausted" + + +def test_a_failed_transcription_reports_the_providers_message(gateway): + """AssemblyAI-backed jobs end as "failed", not "error"; the message must survive.""" + gateway.stt_statuses = ["queued", "failed"] + gateway.stt_error = {"name": "ProviderError", "message": "Internal server error. Please retry."} + assert tools().transcribe_audio("https://files.example/clip.mp3") == ( + "Failed to transcribe audio: Internal server error. Please retry." + ) + + +def test_a_waiting_job_keeps_polling(gateway): + """The Nova-3 docs example still tests for "waiting", so it counts as in-progress.""" + gateway.stt_statuses = ["queued", "waiting", "completed"] + assert tools().transcribe_audio("https://files.example/clip.mp3") == "hello from agno" + assert paths(gateway).count(("GET", "/v1/stt/stt-1")) == 2 + + +def test_generate_video_stops_on_an_unknown_status(gateway): + gateway.video_statuses = ["queued", "cancelled"] + ["cancelled"] * 50 + result = tools().generate_video("a boat") + assert result.content == "Failed to generate video: job ended with status 'cancelled'" + assert paths(gateway).count(("GET", "/v2/video/generations")) == 1 + + +def test_generate_video_retries_a_transient_poll_error(gateway): + gateway.poll_failures = [503, 429] + result = tools().generate_video("a boat") + assert result.content == "Video generated successfully." + assert paths(gateway).count(("GET", "/v2/video/generations")) == 4 + + +def test_generate_video_gives_up_on_a_persistent_poll_error(gateway): + gateway.poll_failures = [503, 503, 503, 503] + result = tools().generate_video("a boat") + assert result.content == "Failed to generate video: AI/ML API returned HTTP 503: try later" + + +def test_generate_video_gives_up_after_the_timeout(gateway): + gateway.video_statuses = ["queued"] * 50 + with patch("agno.tools.models.aimlapi.time.monotonic", side_effect=[0, 0, 1000]): + result = tools(video_timeout=10).generate_video("a boat") + assert result.content == "Failed to generate video: video generation still running after 10s" + + +# --- generate_speech ----------------------------------------------------------- + + +def test_generate_speech_returns_audio(gateway): + result = tools(speech_voice="nova", speech_speed=1.2).generate_speech("hello") + assert result.content.startswith("Speech generated successfully with ID: ") + audio = result.audios[0] + assert (audio.content, audio.mime_type, audio.format) == (b"\x00mp3", "audio/mpeg", "mp3") + assert json.loads(gateway.calls[0].content) == { + "model": "openai/tts-1", + "text": "hello", + "response_format": "mp3", + "voice": "nova", + "speed": 1.2, + } + + +@pytest.mark.asyncio +async def test_agenerate_speech_returns_audio(gateway): + result = await tools().agenerate_speech("hello") + assert result.audios[0].format == "mp3" + + +# --- transcribe_audio ---------------------------------------------------------- + + +def test_transcribe_audio_uploads_a_file_from_the_base_dir(gateway, tmp_path): + (tmp_path / "clip.mp3").write_bytes(b"\x00mp3") + assert tools(base_dir=tmp_path).transcribe_audio("clip.mp3") == "hello from agno" + submit = gateway.calls[0] + assert submit.url.path == "/v1/stt/create" + assert submit.headers["content-type"].startswith("multipart/form-data") + assert b'name="model"' in submit.content and b"deepgram/nova-3" in submit.content + assert b'filename="clip.mp3"' in submit.content + assert paths(gateway)[1:] == [("GET", "/v1/stt/stt-1")] + + +@pytest.mark.asyncio +async def test_atranscribe_audio_uploads_a_file(gateway, tmp_path): + (tmp_path / "clip.mp3").write_bytes(b"\x00mp3") + assert await tools(base_dir=tmp_path).atranscribe_audio("clip.mp3") == "hello from agno" + assert paths(gateway) == [("POST", "/v1/stt/create"), ("GET", "/v1/stt/stt-1")] + + +def test_transcribe_audio_refuses_paths_outside_the_base_dir(gateway, tmp_path): + secret = tmp_path / "secret.key" + secret.write_bytes(b"private") + sandbox = tmp_path / "audio" + sandbox.mkdir() + result = tools(base_dir=sandbox).transcribe_audio("../secret.key") + assert result.startswith("Failed to transcribe audio: ") + assert "outside the allowed directory" in result + assert gateway.calls == [] + + +def test_transcribe_audio_passes_a_url_through(gateway): + assert tools(transcription_language="en").transcribe_audio("https://files.example/clip.mp3") == "hello from agno" + assert json.loads(gateway.calls[0].content) == { + "model": "deepgram/nova-3", + "language": "en", + "url": "https://files.example/clip.mp3", + } + + +def test_an_empty_transcript_is_a_transcript(gateway): + gateway.transcript = "" + assert tools().transcribe_audio("https://files.example/silence.mp3") == "" + + +def test_transcribe_audio_reports_a_missing_file(gateway, tmp_path): + assert tools(base_dir=tmp_path).transcribe_audio("nowhere.mp3").startswith("Failed to transcribe audio: ") + assert gateway.calls == [] diff --git a/libs/agno/tests/unit/tools/test_coding_tools.py b/libs/agno/tests/unit/tools/test_coding_tools.py index 8b18e18a89e..209280e3319 100644 --- a/libs/agno/tests/unit/tools/test_coding_tools.py +++ b/libs/agno/tests/unit/tools/test_coding_tools.py @@ -506,7 +506,7 @@ def test_enable_flags(): """Test that tools can be individually disabled.""" with tempfile.TemporaryDirectory() as tmp_dir: base_dir = Path(tmp_dir) - tools = CodingTools(base_dir=base_dir, enable_read_file=False) + tools = CodingTools(base_dir=base_dir, enable_read_file=False, enable_run_shell=True) tool_names = [fn for fn in tools.functions] assert "read_file" not in tool_names @@ -515,18 +515,18 @@ def test_enable_flags(): assert "run_shell" in tool_names -def test_exploration_tools_disabled_by_default(): - """Test that grep, find, ls are disabled by default.""" +def test_optional_tools_disabled_by_default(): + """Test that run_shell, grep, find, ls are disabled by default.""" with tempfile.TemporaryDirectory() as tmp_dir: base_dir = Path(tmp_dir) tools = CodingTools(base_dir=base_dir) tool_names = list(tools.functions.keys()) - assert len(tool_names) == 4 + assert len(tool_names) == 3 assert "read_file" in tool_names assert "edit_file" in tool_names assert "write_file" in tool_names - assert "run_shell" in tool_names + assert "run_shell" not in tool_names assert "grep" not in tool_names assert "find" not in tool_names assert "ls" not in tool_names @@ -552,51 +552,86 @@ def test_all_flag(): # --- shell sandbox tests --- -def test_run_shell_blocks_metacharacters(): - """Test that shell metacharacters are blocked in restricted mode.""" +def test_run_shell_rejects_control_operators(): + """Standalone control operators are rejected as unsupported in restricted mode. + + Restricted mode runs without a shell, so these cannot chain anyway; rejecting the + common spaced form gives a clear error instead of a confusing literal run. + """ with tempfile.TemporaryDirectory() as tmp_dir: base_dir = Path(tmp_dir) tools = CodingTools(base_dir=base_dir) - # Command chaining with && - result = tools.run_shell("echo hello && cat /etc/passwd") - assert "Error" in result - assert "&&" in result + for op, cmd in ( + ("&&", "echo hello && cat /etc/passwd"), + ("||", "false || cat /etc/passwd"), + (";", "echo hello ; cat /etc/passwd"), + ("|", "echo hello | cat"), + ("&", "echo hello & echo pwned"), + (">", "echo hello > escaped.txt"), + (">>", "echo hello >> escaped.txt"), + ("<", "cat < /etc/passwd"), + ): + result = tools.run_shell(cmd) + assert "Error" in result + assert "not supported in restricted mode" in result + assert op in result - # Command chaining with || - result = tools.run_shell("false || cat /etc/passwd") - assert "Error" in result - assert "||" in result - # Command chaining with ; - result = tools.run_shell("echo hello; cat /etc/passwd") - assert "Error" in result - assert ";" in result +def test_run_shell_operators_are_inert_without_a_shell(): + """Chaining, substitution, redirection, and globbing must not execute. - # Pipe - result = tools.run_shell("echo hello | cat") - assert "Error" in result - assert "|" in result + Restricted mode runs the tokenized command directly (shell=False), so operators + are passed as literal arguments. Proven by side effects that never happen. + """ + with tempfile.TemporaryDirectory() as tmp_dir: + base_dir = Path(tmp_dir) + tools = CodingTools(base_dir=base_dir) - # Command substitution with $() - result = tools.run_shell("echo $(cat /etc/passwd)") - assert "Error" in result - assert "$(" in result + # Command substitution does not run: touch is never executed. + for cmd in ( + "echo $(touch pwned.txt)", + 'echo "$(touch pwned.txt)"', + "echo `touch pwned.txt`", + ): + result = tools.run_shell(cmd) + assert not (base_dir / "pwned.txt").exists() - # Command substitution with backticks - result = tools.run_shell("echo `cat /etc/passwd`") - assert "Error" in result - assert "`" in result + # Glued chaining does not run a second command: only echo executes. + result = tools.run_shell("echo hello;touch chained.txt") + assert not (base_dir / "chained.txt").exists() - # Output redirection - result = tools.run_shell("echo hello > /tmp/evil.txt") - assert "Error" in result - assert ">" in result + # Globs are not expanded: the literal pattern is passed through. + (base_dir / "a.py").write_text("x\n") + result = tools.run_shell("echo *.py") + assert "*.py" in result - # Input redirection - result = tools.run_shell("cat < /etc/passwd") - assert "Error" in result - assert "<" in result + +def test_run_shell_allows_quoted_operators(): + """Operators inside quotes are ordinary characters and must not be rejected. + + These are valid commands (e.g. a commit message containing '&') that a raw + substring check wrongly blocked. + """ + with tempfile.TemporaryDirectory() as tmp_dir: + base_dir = Path(tmp_dir) + tools = CodingTools(base_dir=base_dir) + + for cmd in ( + "echo 'A & B'", # single-quoted separator + 'echo "A & B"', # double-quoted separator + "echo 'a;b'", + "echo 'a|b'", + "echo 'a>b'", + "echo 'a str: + return "first" + + @server.tool + def second_tool() -> str: + return "second" + + async with Client(server) as client: + session = client.session + assert isinstance(session, ClientSession) + toolkit = MCPTools(session=session) + with patch.object(session, "list_tools", wraps=session.list_tools) as list_tools: + await toolkit.build_tools() + assert list_tools.await_count == 2 + assert set(toolkit.get_async_functions()) == {"first_tool", "second_tool"} + result = await toolkit.functions["second_tool"].entrypoint() + + assert result.content == "second" diff --git a/libs/agno/tests/unit/tools/test_python_tools.py b/libs/agno/tests/unit/tools/test_python_tools.py index 7253026fe64..e27885d6b74 100644 --- a/libs/agno/tests/unit/tools/test_python_tools.py +++ b/libs/agno/tests/unit/tools/test_python_tools.py @@ -204,3 +204,47 @@ def test_run_python_file_blocks_path_traversal(temp_dir): result = python_tools.run_python_file_return_variable("../malicious.py") assert "outside the allowed base directory" in result + + +# restrict_to_base_dir does not sandbox executed code — pin the documented limitation +# so nobody later mistakes the path-traversal guards above for a code sandbox. +def test_run_python_code_ignores_restrict_to_base_dir(temp_dir): + """run_python_code executes regardless of restrict_to_base_dir: it can read outside base_dir. + + The path-traversal guards only cover file-path arguments to the file helpers. + Executed code goes straight to exec(), so restrict_to_base_dir is not a sandbox. + """ + outside = temp_dir.parent / "outside_secret.txt" + outside.write_text("top-secret") + try: + python_tools = PythonTools(base_dir=temp_dir / "sandbox", restrict_to_base_dir=True) + code = f"data = open({str(outside)!r}).read()" + result = python_tools.run_python_code(code, "data") + assert result == "top-secret" + finally: + outside.unlink(missing_ok=True) + + +def test_requires_confirmation_gates_execution_tools(temp_dir): + """The documented mitigation works: requires_confirmation_tools marks exec tools for HITL approval.""" + python_tools = PythonTools( + base_dir=temp_dir, + requires_confirmation_tools=["run_python_code", "save_to_file_and_run"], + ) + assert python_tools.functions["run_python_code"].requires_confirmation is True + assert python_tools.functions["save_to_file_and_run"].requires_confirmation is True + + +def test_exclude_tools_drops_execution_tools(temp_dir): + """The documented mitigation works: exclude_tools removes the code-execution entry points.""" + python_tools = PythonTools( + base_dir=temp_dir, + exclude_tools=["run_python_code", "save_to_file_and_run", "run_python_file_return_variable"], + ) + registered = set(python_tools.functions.keys()) + assert "run_python_code" not in registered + assert "save_to_file_and_run" not in registered + assert "run_python_file_return_variable" not in registered + # Benign helpers remain available. + assert "read_file" in registered + assert "list_files" in registered diff --git a/libs/agno/tests/unit/vectordb/test_pineconedb.py b/libs/agno/tests/unit/vectordb/test_pineconedb.py index 42b7a7b2143..b9da44ced78 100644 --- a/libs/agno/tests/unit/vectordb/test_pineconedb.py +++ b/libs/agno/tests/unit/vectordb/test_pineconedb.py @@ -263,7 +263,7 @@ def test_search(mock_pinecone_db, mock_embedder): # Check that index.query was called with the right arguments mock_pinecone_db.index.query.assert_called_with( - vector=[0.1] * 1024, top_k=2, namespace=TEST_NAMESPACE, filter=None, include_values=None, include_metadata=True + vector=[0.1] * 1024, top_k=2, namespace=TEST_NAMESPACE, filter=None, include_values=False, include_metadata=True ) # Check the results diff --git a/libs/agno/tests/unit/workflow/test_previous_content_values.py b/libs/agno/tests/unit/workflow/test_previous_content_values.py new file mode 100644 index 00000000000..442553ae46a --- /dev/null +++ b/libs/agno/tests/unit/workflow/test_previous_content_values.py @@ -0,0 +1,102 @@ +"""Unit tests for previous step content that is falsy but not empty (0, False, empty containers). + +Content that is None or blank text is still skipped. +""" + +import pytest + +from agno.workflow.parallel import Parallel +from agno.workflow.step import Step +from agno.workflow.types import StepInput, StepOutput, StepType + + +def summarize(step_input: StepInput) -> StepOutput: + return StepOutput(content="") + + +@pytest.mark.parametrize("content", [0, 0.0, False, [], {}]) +def test_previous_content_preserves_falsy_values(content): + """get_all_previous_content keeps a step whose content is falsy but not empty.""" + step_input = StepInput(previous_step_outputs={"result": StepOutput(content=content)}) + + assert step_input.get_all_previous_content() == f"=== result ===\n{content}" + + +def test_previous_content_skips_none_and_blank_text_and_preserves_order(): + """None and blank text are skipped, and the remaining steps keep their order.""" + step_input = StepInput( + previous_step_outputs={ + "count": StepOutput(content=0), + "missing": StepOutput(content=None), + "blank": StepOutput(content=" "), + "approved": StepOutput(content=False), + } + ) + + assert step_input.get_all_previous_content() == "=== count ===\n0\n\n=== approved ===\nFalse" + + +def test_parallel_step_content_preserves_falsy_values(): + """get_step_content on a Parallel step keeps falsy sub-step content and skips blank text.""" + parallel_output = StepOutput( + step_name="parallel", + step_type=StepType.PARALLEL, + steps=[ + StepOutput(step_name="count", content=0), + StepOutput(step_name="approved", content=False), + StepOutput(step_name="blank", content=" "), + ], + ) + step_input = StepInput(previous_step_outputs={"parallel": parallel_output}) + + assert step_input.get_step_content("parallel") == {"count": "0", "approved": "False"} + + +def test_parallel_nested_step_content_preserves_falsy_values(): + """get_step_content on a Parallel keeps falsy content from steps nested inside a sub-step such as a Condition.""" + parallel_output = StepOutput( + step_name="parallel", + step_type=StepType.PARALLEL, + steps=[ + StepOutput( + step_name="condition", + step_type=StepType.CONDITION, + content="Condition completed", + steps=[StepOutput(step_name="count", content=0), StepOutput(step_name="blank", content=" ")], + ), + StepOutput(step_name="label", content="ok"), + ], + ) + step_input = StepInput(previous_step_outputs={"parallel": parallel_output}) + + assert step_input.get_step_content("parallel") == {"count": "0", "label": "ok"} + + +def test_parallel_aggregated_content_preserves_falsy_values(): + """Parallel aggregated content shows falsy content instead of *(No content)*.""" + parallel = Parallel(name="parallel") + + content = parallel._build_aggregated_content( + [StepOutput(step_name="count", content=0), StepOutput(step_name="approved", content=False)] + ) + + assert "count\n0\n" in content + assert "approved\nFalse" in content + assert "*(No content)*" not in content + + +def test_next_step_input_after_parallel_preserves_falsy_values(): + """The input built for the step after a Parallel keeps falsy sub-step content and skips blank text.""" + step = Step(name="summary", executor=summarize) + parallel_output = StepOutput( + step_name="parallel", + step_type=StepType.PARALLEL, + content="aggregated", + steps=[ + StepOutput(step_name="count", content=0), + StepOutput(step_name="approved", content=False), + StepOutput(step_name="blank", content=" "), + ], + ) + + assert step._get_deepest_content_from_step_output(parallel_output) == "=== count ===\n0\n\n=== approved ===\nFalse" diff --git a/libs/agno_infra/README.md b/libs/agno_infra/README.md index 3af4dc1b646..dd0ca72da5d 100644 --- a/libs/agno_infra/README.md +++ b/libs/agno_infra/README.md @@ -126,7 +126,7 @@ agno/ ## 📄 License -This project is licensed under the Mozilla Public License 2.0 - see the [LICENSE](LICENSE) file for details. +This project is licensed under the Apache-2.0 license - see the [LICENSE](LICENSE) file for details. ## 🙋‍♀️ Support diff --git a/libs/agnoctl/agnoctl/clients/base.py b/libs/agnoctl/agnoctl/clients/base.py index 1617d127ba9..28bfca1d998 100644 --- a/libs/agnoctl/agnoctl/clients/base.py +++ b/libs/agnoctl/agnoctl/clients/base.py @@ -75,7 +75,7 @@ def read_json_lenient(path: Path) -> Optional[Dict[str, Any]]: if not path.exists(): return None try: - parsed = json.loads(path.read_text()) + parsed = json.loads(path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError): return None return parsed if isinstance(parsed, dict) else None @@ -86,7 +86,7 @@ def read_json_strict(path: Path) -> Dict[str, Any]: if not path.exists(): return {} try: - parsed = json.loads(path.read_text()) + parsed = json.loads(path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError) as e: raise CLIError( "Refusing to modify " + str(path) + ": the existing file is not valid JSON (" + str(e) + ").", @@ -127,7 +127,7 @@ def atomic_write_text(path: Path, text: str, *, secure: bool) -> None: fd, tmp_name = tempfile.mkstemp(dir=str(path.parent), prefix="." + path.name + ".", suffix=".tmp") tmp = Path(tmp_name) try: - with os.fdopen(fd, "w") as handle: + with os.fdopen(fd, "w", encoding="utf-8") as handle: handle.write(text) handle.flush() os.fsync(handle.fileno()) diff --git a/libs/agnoctl/agnoctl/clients/codex.py b/libs/agnoctl/agnoctl/clients/codex.py index dad883c8c81..91a2948def7 100644 --- a/libs/agnoctl/agnoctl/clients/codex.py +++ b/libs/agnoctl/agnoctl/clients/codex.py @@ -167,7 +167,7 @@ def _read_strict(self) -> "tuple[str, Dict[str, Any]]": (shared by the write and remove paths so their refusal behavior cannot drift).""" if not self.config_path.exists(): return "", {} - text = self.config_path.read_text() + text = self.config_path.read_text(encoding="utf-8") try: return text, tomllib.loads(text) except tomllib.TOMLDecodeError as e: @@ -188,7 +188,7 @@ def _parse_config(self) -> Optional[Dict[str, Any]]: if not self.config_path.exists(): return None try: - return tomllib.loads(self.config_path.read_text()) + return tomllib.loads(self.config_path.read_text(encoding="utf-8")) except (OSError, tomllib.TOMLDecodeError): return None diff --git a/libs/agnoctl/pyproject.toml b/libs/agnoctl/pyproject.toml index a0eaae11f2d..dea78f5e97e 100644 --- a/libs/agnoctl/pyproject.toml +++ b/libs/agnoctl/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "agnoctl" -version = "0.2.0" +version = "0.2.1" description = "The Agno CLI: connect and operate AgentOS from the terminal, built for humans and coding agents" requires-python = ">=3.9,<4" readme = "README.md" diff --git a/libs/agnoctl/tests/test_adapters.py b/libs/agnoctl/tests/test_adapters.py index b2bbc8e489d..9fd0079f498 100644 --- a/libs/agnoctl/tests/test_adapters.py +++ b/libs/agnoctl/tests/test_adapters.py @@ -543,6 +543,42 @@ def test_atomic_write_text_direct_secure(tmp_path: Path, permissive_umask): assert _mode(target) == 0o600 +@pytest.fixture +def cp1252_default_encoding(monkeypatch): + """Run the test as if on a Windows machine whose default text encoding is cp1252, so a + config read or write that does not name an encoding would mangle non-ASCII text.""" + read_text = Path.read_text + fdopen = os.fdopen + + def read_text_cp1252(self, encoding=None, **kwargs): + return read_text(self, encoding=encoding or "cp1252", **kwargs) + + def fdopen_cp1252(fd, mode="r", *args, encoding=None, **kwargs): + if encoding is None and "b" not in mode: + encoding = "cp1252" + return fdopen(fd, mode, *args, encoding=encoding, **kwargs) + + monkeypatch.setattr(Path, "read_text", read_text_cp1252) + monkeypatch.setattr(base_module.os, "fdopen", fdopen_cp1252) + + +def test_writes_keep_non_ascii_config_text(tmp_path: Path, cp1252_default_encoding): + """Client configs are UTF-8 on disk. Merging an entry must not re-encode the user's + existing non-ASCII text through the platform's default encoding.""" + for adapter, path in _file_writing_adapters(tmp_path): + seed = "# José\n" if path.suffix == ".toml" else json.dumps({"note": "José"}, ensure_ascii=False) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(seed.encode("utf-8")) + + adapter.write("agno", URL, TOKEN) + + text = path.read_bytes().decode("utf-8") + if path.suffix == ".toml": + assert text.startswith("# José\n"), adapter.key + else: + assert json.loads(text)["note"] == "José", adapter.key + + # -- remove ------------------------------------------------------------------------------