Merge pull request #5 from ggozad/feat/qa
Question/Answering using LLMs
This commit is contained in:
commit
e5d8396d1b
17 changed files with 332 additions and 54 deletions
|
|
@ -7,3 +7,7 @@ We use [repliqa](https://huggingface.co/datasets/ServiceNow/repliqa) for the eva
|
||||||
We load the `News Stories` from `repliqa_3` which is 1035 documents, using `tests/generate_benchmark_db.py`, using the `mxbai-embed-large` Ollama embeddings.
|
We load the `News Stories` from `repliqa_3` which is 1035 documents, using `tests/generate_benchmark_db.py`, using the `mxbai-embed-large` Ollama embeddings.
|
||||||
|
|
||||||
Subsequently, we run a search over the `question` for each row of the dataset and check whether we match the document that answers the question. The recall obtained is ~0.75 for matching in the top result, raising to ~0.75 for the top 3 results.
|
Subsequently, we run a search over the `question` for each row of the dataset and check whether we match the document that answers the question. The recall obtained is ~0.75 for matching in the top result, raising to ~0.75 for the top 3 results.
|
||||||
|
|
||||||
|
* Question/Answer evaluation
|
||||||
|
|
||||||
|
We use the `News Stories` from `repliqa_3` using the `mxbai-embed-large` Ollama embeddings, with a QA agent also using Ollama with the `qwen3` model (8b). For each story we ask the `question` and use an LLM judge (also `qwen3`) to evaluate whether the answer is correct or not. Thus we obtain accuracy of ~0.54.
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ Retrieval-Augmented Generation (RAG) library on SQLite.
|
||||||
- **Local SQLite**: No external servers required
|
- **Local SQLite**: No external servers required
|
||||||
- **Multiple embedding providers**: Ollama, VoyageAI, OpenAI
|
- **Multiple embedding providers**: Ollama, VoyageAI, OpenAI
|
||||||
- **Hybrid search**: Vector + full-text search with Reciprocal Rank Fusion
|
- **Hybrid search**: Vector + full-text search with Reciprocal Rank Fusion
|
||||||
|
- **Question answering**: Built-in QA agents on your documents
|
||||||
- **File monitoring**: Auto-index files when run as server
|
- **File monitoring**: Auto-index files when run as server
|
||||||
- **40+ file formats**: PDF, DOCX, HTML, Markdown, audio, URLs
|
- **40+ file formats**: PDF, DOCX, HTML, Markdown, audio, URLs
|
||||||
- **MCP server**: Expose as tools for AI assistants
|
- **MCP server**: Expose as tools for AI assistants
|
||||||
|
|
@ -27,6 +28,9 @@ haiku-rag add-src document.pdf
|
||||||
# Search
|
# Search
|
||||||
haiku-rag search "query"
|
haiku-rag search "query"
|
||||||
|
|
||||||
|
# Ask questions
|
||||||
|
haiku-rag ask "Who is the author of haiku.rag?"
|
||||||
|
|
||||||
# Start server with file monitoring
|
# Start server with file monitoring
|
||||||
export MONITOR_DIRECTORIES="/path/to/docs"
|
export MONITOR_DIRECTORIES="/path/to/docs"
|
||||||
haiku-rag serve
|
haiku-rag serve
|
||||||
|
|
@ -45,6 +49,10 @@ async with HaikuRAG("database.db") as client:
|
||||||
results = await client.search("query")
|
results = await client.search("query")
|
||||||
for chunk, score in results:
|
for chunk, score in results:
|
||||||
print(f"{score:.3f}: {chunk.content}")
|
print(f"{score:.3f}: {chunk.content}")
|
||||||
|
|
||||||
|
# Ask questions
|
||||||
|
answer = await client.ask("Who is the author of haiku.rag?")
|
||||||
|
print(answer)
|
||||||
```
|
```
|
||||||
|
|
||||||
## MCP Server
|
## MCP Server
|
||||||
|
|
|
||||||
|
|
@ -47,6 +47,15 @@ With options:
|
||||||
haiku-rag search "python programming" --limit 10 --k 100
|
haiku-rag search "python programming" --limit 10 --k 100
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Question Answering
|
||||||
|
|
||||||
|
Ask questions about your documents:
|
||||||
|
```bash
|
||||||
|
haiku-rag ask "Who is the author of haiku.rag?"
|
||||||
|
```
|
||||||
|
|
||||||
|
The QA agent will search your documents for relevant information and provide a comprehensive answer.
|
||||||
|
|
||||||
## Server
|
## Server
|
||||||
|
|
||||||
Start the MCP server:
|
Start the MCP server:
|
||||||
|
|
|
||||||
|
|
@ -55,3 +55,50 @@ EMBEDDINGS_MODEL="text-embedding-3-small" # or text-embedding-3-large
|
||||||
EMBEDDINGS_VECTOR_DIM=1536
|
EMBEDDINGS_VECTOR_DIM=1536
|
||||||
OPENAI_API_KEY="your-api-key"
|
OPENAI_API_KEY="your-api-key"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Question Answering Providers
|
||||||
|
|
||||||
|
Configure which LLM provider to use for question answering.
|
||||||
|
|
||||||
|
### Ollama (Default)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
QA_PROVIDER="ollama"
|
||||||
|
QA_MODEL="qwen3"
|
||||||
|
OLLAMA_BASE_URL="http://localhost:11434"
|
||||||
|
```
|
||||||
|
|
||||||
|
### OpenAI
|
||||||
|
|
||||||
|
For OpenAI QA, you need to install haiku.rag with OpenAI extras:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv pip install haiku.rag --extra openai
|
||||||
|
```
|
||||||
|
|
||||||
|
Then configure:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
QA_PROVIDER="openai"
|
||||||
|
QA_MODEL="gpt-4o-mini" # or gpt-4, gpt-3.5-turbo, etc.
|
||||||
|
OPENAI_API_KEY="your-api-key"
|
||||||
|
```
|
||||||
|
|
||||||
|
## Other Settings
|
||||||
|
|
||||||
|
### Database and Storage
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Default data directory (where SQLite database is stored)
|
||||||
|
DEFAULT_DATA_DIR="/path/to/data"
|
||||||
|
```
|
||||||
|
|
||||||
|
### Document Processing
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Chunk size for document processing
|
||||||
|
CHUNK_SIZE=256
|
||||||
|
|
||||||
|
# Chunk overlap for better context
|
||||||
|
CHUNK_OVERLAP=32
|
||||||
|
```
|
||||||
|
|
|
||||||
|
|
@ -8,6 +8,7 @@
|
||||||
- **Local SQLite**: No need to run additional servers
|
- **Local SQLite**: No need to run additional servers
|
||||||
- **Support for various embedding providers**: Ollama, VoyageAI, OpenAI or add your own
|
- **Support for various embedding providers**: Ollama, VoyageAI, OpenAI or add your own
|
||||||
- **Hybrid Search**: Vector search using `sqlite-vec` combined with full-text search `FTS5`, using Reciprocal Rank Fusion
|
- **Hybrid Search**: Vector search using `sqlite-vec` combined with full-text search `FTS5`, using Reciprocal Rank Fusion
|
||||||
|
- **Question Answering**: Built-in QA agents using Ollama or OpenAI.
|
||||||
- **File monitoring**: Automatically index files when run as a server
|
- **File monitoring**: Automatically index files when run as a server
|
||||||
- **Extended file format support**: Parse 40+ file formats including PDF, DOCX, HTML, Markdown, audio and more. Or add a URL!
|
- **Extended file format support**: Parse 40+ file formats including PDF, DOCX, HTML, Markdown, audio and more. Or add a URL!
|
||||||
- **MCP server**: Exposes functionality as MCP tools
|
- **MCP server**: Exposes functionality as MCP tools
|
||||||
|
|
@ -31,12 +32,16 @@ async with HaikuRAG("database.db") as client:
|
||||||
|
|
||||||
# Search documents
|
# Search documents
|
||||||
results = await client.search("query")
|
results = await client.search("query")
|
||||||
|
|
||||||
|
# Ask questions
|
||||||
|
answer = await client.ask("Who is the author of haiku.rag?")
|
||||||
```
|
```
|
||||||
|
|
||||||
Or use the CLI:
|
Or use the CLI:
|
||||||
```bash
|
```bash
|
||||||
haiku-rag add "Your document content"
|
haiku-rag add "Your document content"
|
||||||
haiku-rag search "query"
|
haiku-rag search "query"
|
||||||
|
haiku-rag ask "Who is the author of haiku.rag?"
|
||||||
```
|
```
|
||||||
|
|
||||||
## Documentation
|
## Documentation
|
||||||
|
|
@ -44,6 +49,7 @@ haiku-rag search "query"
|
||||||
- [Installation](installation.md) - Install haiku.rag with different providers
|
- [Installation](installation.md) - Install haiku.rag with different providers
|
||||||
- [Configuration](configuration.md) - Environment variables and settings
|
- [Configuration](configuration.md) - Environment variables and settings
|
||||||
- [CLI](cli.md) - Command line interface usage
|
- [CLI](cli.md) - Command line interface usage
|
||||||
|
- [Question Answering](qa.md) - QA agents and natural language queries
|
||||||
- [Server](server.md) - File monitoring and server mode
|
- [Server](server.md) - File monitoring and server mode
|
||||||
- [MCP](mcp.md) - Model Context Protocol integration
|
- [MCP](mcp.md) - Model Context Protocol integration
|
||||||
- [Python](python.md) - Python API reference
|
- [Python](python.md) - Python API reference
|
||||||
|
|
|
||||||
|
|
@ -91,4 +91,19 @@ for chunk, relevance_score in results:
|
||||||
print(f"Relevance: {relevance_score:.3f}")
|
print(f"Relevance: {relevance_score:.3f}")
|
||||||
print(f"Content: {chunk.content}")
|
print(f"Content: {chunk.content}")
|
||||||
print(f"From document: {chunk.document_id}")
|
print(f"From document: {chunk.document_id}")
|
||||||
|
print(f"Document URI: {chunk.document_uri}")
|
||||||
|
print(f"Document metadata: {chunk.document_meta}")
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Question Answering
|
||||||
|
|
||||||
|
Ask questions about your documents:
|
||||||
|
|
||||||
|
```python
|
||||||
|
answer = await client.ask("Who is the author of haiku.rag?")
|
||||||
|
print(answer)
|
||||||
|
```
|
||||||
|
|
||||||
|
The QA agent will search your documents for relevant information and use the configured LLM to generate a comprehensive answer.
|
||||||
|
|
||||||
|
The QA provider and model can be configured via environment variables (see [Configuration](configuration.md)).
|
||||||
|
|
|
||||||
|
|
@ -61,6 +61,17 @@ class HaikuRAGApp:
|
||||||
for chunk, score in results:
|
for chunk, score in results:
|
||||||
self._rich_print_search_result(chunk, score)
|
self._rich_print_search_result(chunk, score)
|
||||||
|
|
||||||
|
async def ask(self, question: str):
|
||||||
|
async with HaikuRAG(db_path=self.db_path) as self.client:
|
||||||
|
try:
|
||||||
|
answer = await self.client.ask(question)
|
||||||
|
self.console.print(f"[bold blue]Question:[/bold blue] {question}")
|
||||||
|
self.console.print()
|
||||||
|
self.console.print("[bold green]Answer:[/bold green]")
|
||||||
|
self.console.print(Markdown(answer))
|
||||||
|
except Exception as e:
|
||||||
|
self.console.print(f"[red]Error: {e}[/red]")
|
||||||
|
|
||||||
def _rich_print_document(self, doc: Document, truncate: bool = False):
|
def _rich_print_document(self, doc: Document, truncate: bool = False):
|
||||||
"""Format a document for display."""
|
"""Format a document for display."""
|
||||||
if truncate:
|
if truncate:
|
||||||
|
|
|
||||||
|
|
@ -113,6 +113,21 @@ def search(
|
||||||
event_loop.run_until_complete(app.search(query=query, limit=limit, k=k))
|
event_loop.run_until_complete(app.search(query=query, limit=limit, k=k))
|
||||||
|
|
||||||
|
|
||||||
|
@cli.command("ask", help="Ask a question using the QA agent")
|
||||||
|
def ask(
|
||||||
|
question: str = typer.Argument(
|
||||||
|
help="The question to ask",
|
||||||
|
),
|
||||||
|
db: Path = typer.Option(
|
||||||
|
get_default_data_dir() / "haiku.rag.sqlite",
|
||||||
|
"--db",
|
||||||
|
help="Path to the SQLite database file",
|
||||||
|
),
|
||||||
|
):
|
||||||
|
app = HaikuRAGApp(db_path=db)
|
||||||
|
event_loop.run_until_complete(app.ask(question=question))
|
||||||
|
|
||||||
|
|
||||||
@cli.command(
|
@cli.command(
|
||||||
"serve", help="Start the haiku.rag MCP server (by default in streamable HTTP mode)"
|
"serve", help="Start the haiku.rag MCP server (by default in streamable HTTP mode)"
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -36,7 +36,7 @@ class HaikuRAG:
|
||||||
"""Async context manager entry."""
|
"""Async context manager entry."""
|
||||||
return self
|
return self
|
||||||
|
|
||||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002
|
||||||
"""Async context manager exit."""
|
"""Async context manager exit."""
|
||||||
self.close()
|
self.close()
|
||||||
return False
|
return False
|
||||||
|
|
@ -256,6 +256,20 @@ class HaikuRAG:
|
||||||
"""
|
"""
|
||||||
return await self.chunk_repository.search_chunks_hybrid(query, limit, k)
|
return await self.chunk_repository.search_chunks_hybrid(query, limit, k)
|
||||||
|
|
||||||
|
async def ask(self, question: str) -> str:
|
||||||
|
"""Ask a question using the configured QA agent.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
question: The question to ask
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The generated answer as a string
|
||||||
|
"""
|
||||||
|
from haiku.rag.qa import get_qa_agent
|
||||||
|
|
||||||
|
qa_agent = get_qa_agent(self)
|
||||||
|
return await qa_agent.answer(question)
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
"""Close the underlying store connection."""
|
"""Close the underlying store connection."""
|
||||||
self.store.close()
|
self.store.close()
|
||||||
|
|
|
||||||
|
|
@ -27,6 +27,10 @@ class AppConfig(BaseModel):
|
||||||
|
|
||||||
OLLAMA_BASE_URL: str = "http://localhost:11434"
|
OLLAMA_BASE_URL: str = "http://localhost:11434"
|
||||||
|
|
||||||
|
# Provider keys
|
||||||
|
VOYAGE_API_KEY: str = ""
|
||||||
|
OPENAI_API_KEY: str = ""
|
||||||
|
|
||||||
@field_validator("MONITOR_DIRECTORIES", mode="before")
|
@field_validator("MONITOR_DIRECTORIES", mode="before")
|
||||||
@classmethod
|
@classmethod
|
||||||
def parse_monitor_directories(cls, v):
|
def parse_monitor_directories(cls, v):
|
||||||
|
|
@ -41,3 +45,7 @@ class AppConfig(BaseModel):
|
||||||
|
|
||||||
# Expose Config object for app to import
|
# Expose Config object for app to import
|
||||||
Config = AppConfig.model_validate(os.environ)
|
Config = AppConfig.model_validate(os.environ)
|
||||||
|
if Config.OPENAI_API_KEY:
|
||||||
|
os.environ["OPENAI_API_KEY"] = Config.OPENAI_API_KEY
|
||||||
|
if Config.VOYAGE_API_KEY:
|
||||||
|
os.environ["VOYAGE_API_KEY"] = Config.VOYAGE_API_KEY
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,26 @@
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
from haiku.rag.config import Config
|
||||||
|
from haiku.rag.qa.base import QuestionAnswerAgentBase
|
||||||
|
from haiku.rag.qa.ollama import QuestionAnswerOllamaAgent
|
||||||
|
|
||||||
|
|
||||||
|
def get_qa_agent(client: HaikuRAG, model: str = "") -> QuestionAnswerAgentBase:
|
||||||
|
"""
|
||||||
|
Factory function to get the appropriate QA agent based on the configuration.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if Config.QA_PROVIDER == "ollama":
|
||||||
|
return QuestionAnswerOllamaAgent(client, model or Config.QA_MODEL)
|
||||||
|
|
||||||
|
if Config.QA_PROVIDER == "openai":
|
||||||
|
try:
|
||||||
|
from haiku.rag.qa.openai import QuestionAnswerOpenAIAgent
|
||||||
|
except ImportError:
|
||||||
|
raise ImportError(
|
||||||
|
"OpenAI QA agent requires the 'openai' package. "
|
||||||
|
"Please install haiku.rag with the 'openai' extra:"
|
||||||
|
"uv pip install haiku.rag --extra openai"
|
||||||
|
)
|
||||||
|
return QuestionAnswerOpenAIAgent(client, model or "gpt-4o-mini")
|
||||||
|
|
||||||
|
raise ValueError(f"Unsupported QA provider: {Config.QA_PROVIDER}")
|
||||||
|
|
@ -2,7 +2,7 @@ from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.qa.prompts import SYSTEM_PROMPT
|
from haiku.rag.qa.prompts import SYSTEM_PROMPT
|
||||||
|
|
||||||
|
|
||||||
class QABase:
|
class QuestionAnswerAgentBase:
|
||||||
_model: str = ""
|
_model: str = ""
|
||||||
_system_prompt: str = SYSTEM_PROMPT
|
_system_prompt: str = SYSTEM_PROMPT
|
||||||
|
|
||||||
|
|
@ -14,3 +14,28 @@ class QABase:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"QABase is an abstract class. Please implement the answer method in a subclass."
|
"QABase is an abstract class. Please implement the answer method in a subclass."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
tools = [
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "search_documents",
|
||||||
|
"description": "Search the knowledge base for relevant documents",
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"query": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "The search query to find relevant documents",
|
||||||
|
},
|
||||||
|
"limit": {
|
||||||
|
"type": "integer",
|
||||||
|
"description": "Maximum number of results to return",
|
||||||
|
"default": 3,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": ["query"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
|
||||||
|
|
@ -2,12 +2,12 @@ from ollama import AsyncClient
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
from haiku.rag.qa.base import QABase
|
from haiku.rag.qa.base import QuestionAnswerAgentBase
|
||||||
|
|
||||||
OLLAMA_OPTIONS = {"temperature": 0.0, "seed": 42, "num_ctx": 64000}
|
OLLAMA_OPTIONS = {"temperature": 0.0, "seed": 42, "num_ctx": 64000}
|
||||||
|
|
||||||
|
|
||||||
class QA(QABase):
|
class QuestionAnswerOllamaAgent(QuestionAnswerAgentBase):
|
||||||
def __init__(self, client: HaikuRAG, model: str = Config.QA_MODEL):
|
def __init__(self, client: HaikuRAG, model: str = Config.QA_MODEL):
|
||||||
super().__init__(client, model or self._model)
|
super().__init__(client, model or self._model)
|
||||||
|
|
||||||
|
|
@ -15,30 +15,6 @@ class QA(QABase):
|
||||||
ollama_client = AsyncClient(host=Config.OLLAMA_BASE_URL)
|
ollama_client = AsyncClient(host=Config.OLLAMA_BASE_URL)
|
||||||
|
|
||||||
# Define the search tool
|
# Define the search tool
|
||||||
tools = [
|
|
||||||
{
|
|
||||||
"type": "function",
|
|
||||||
"function": {
|
|
||||||
"name": "search_documents",
|
|
||||||
"description": "Search the knowledge base for relevant documents",
|
|
||||||
"parameters": {
|
|
||||||
"type": "object",
|
|
||||||
"properties": {
|
|
||||||
"query": {
|
|
||||||
"type": "string",
|
|
||||||
"description": "The search query to find relevant documents",
|
|
||||||
},
|
|
||||||
"limit": {
|
|
||||||
"type": "integer",
|
|
||||||
"description": "Maximum number of results to return",
|
|
||||||
"default": 3,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"required": ["query"],
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
]
|
|
||||||
|
|
||||||
messages = [
|
messages = [
|
||||||
{"role": "system", "content": self._system_prompt},
|
{"role": "system", "content": self._system_prompt},
|
||||||
|
|
@ -49,7 +25,7 @@ class QA(QABase):
|
||||||
response = await ollama_client.chat(
|
response = await ollama_client.chat(
|
||||||
model=self._model,
|
model=self._model,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
tools=tools,
|
tools=self.tools,
|
||||||
options=OLLAMA_OPTIONS,
|
options=OLLAMA_OPTIONS,
|
||||||
think=False,
|
think=False,
|
||||||
)
|
)
|
||||||
|
|
|
||||||
101
src/haiku/rag/qa/openai.py
Normal file
101
src/haiku/rag/qa/openai.py
Normal file
|
|
@ -0,0 +1,101 @@
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
try:
|
||||||
|
from openai import AsyncOpenAI
|
||||||
|
from openai.types.chat import (
|
||||||
|
ChatCompletionAssistantMessageParam,
|
||||||
|
ChatCompletionMessageParam,
|
||||||
|
ChatCompletionSystemMessageParam,
|
||||||
|
ChatCompletionToolMessageParam,
|
||||||
|
ChatCompletionUserMessageParam,
|
||||||
|
)
|
||||||
|
from openai.types.chat.chat_completion_tool_param import ChatCompletionToolParam
|
||||||
|
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
from haiku.rag.qa.base import QuestionAnswerAgentBase
|
||||||
|
|
||||||
|
class QuestionAnswerOpenAIAgent(QuestionAnswerAgentBase):
|
||||||
|
def __init__(self, client: HaikuRAG, model: str = "gpt-4o-mini"):
|
||||||
|
super().__init__(client, model or self._model)
|
||||||
|
self.tools: Sequence[ChatCompletionToolParam] = [
|
||||||
|
ChatCompletionToolParam(tool) for tool in self.tools
|
||||||
|
]
|
||||||
|
|
||||||
|
async def answer(self, question: str) -> str:
|
||||||
|
openai_client = AsyncOpenAI()
|
||||||
|
|
||||||
|
# Define the search tool
|
||||||
|
|
||||||
|
messages: list[ChatCompletionMessageParam] = [
|
||||||
|
ChatCompletionSystemMessageParam(
|
||||||
|
role="system", content=self._system_prompt
|
||||||
|
),
|
||||||
|
ChatCompletionUserMessageParam(role="user", content=question),
|
||||||
|
]
|
||||||
|
|
||||||
|
# Initial response with tool calling
|
||||||
|
response = await openai_client.chat.completions.create(
|
||||||
|
model=self._model,
|
||||||
|
messages=messages,
|
||||||
|
tools=self.tools,
|
||||||
|
temperature=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
response_message = response.choices[0].message
|
||||||
|
|
||||||
|
if response_message.tool_calls:
|
||||||
|
messages.append(
|
||||||
|
ChatCompletionAssistantMessageParam(
|
||||||
|
role="assistant",
|
||||||
|
content=response_message.content,
|
||||||
|
tool_calls=[
|
||||||
|
{
|
||||||
|
"id": tc.id,
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": tc.function.name,
|
||||||
|
"arguments": tc.function.arguments,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
for tc in response_message.tool_calls
|
||||||
|
],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
for tool_call in response_message.tool_calls:
|
||||||
|
if tool_call.function.name == "search_documents":
|
||||||
|
import json
|
||||||
|
|
||||||
|
args = json.loads(tool_call.function.arguments)
|
||||||
|
query = args.get("query", question)
|
||||||
|
limit = int(args.get("limit", 3))
|
||||||
|
|
||||||
|
search_results = await self._client.search(query, limit=limit)
|
||||||
|
|
||||||
|
context_chunks = []
|
||||||
|
for chunk, score in search_results:
|
||||||
|
context_chunks.append(
|
||||||
|
f"Content: {chunk.content}\nScore: {score:.4f}"
|
||||||
|
)
|
||||||
|
|
||||||
|
context = "\n\n".join(context_chunks)
|
||||||
|
|
||||||
|
messages.append(
|
||||||
|
ChatCompletionToolMessageParam(
|
||||||
|
role="tool",
|
||||||
|
content=context,
|
||||||
|
tool_call_id=tool_call.id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
final_response = await openai_client.chat.completions.create(
|
||||||
|
model=self._model,
|
||||||
|
messages=messages,
|
||||||
|
temperature=0.0,
|
||||||
|
)
|
||||||
|
return final_response.choices[0].message.content or ""
|
||||||
|
else:
|
||||||
|
return response_message.content or ""
|
||||||
|
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
@ -3,8 +3,6 @@ from pathlib import Path
|
||||||
import pytest
|
import pytest
|
||||||
from datasets import Dataset, load_dataset, load_from_disk
|
from datasets import Dataset, load_dataset, load_from_disk
|
||||||
|
|
||||||
from .llm_judge import LLMJudge
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
@pytest.fixture(scope="session")
|
||||||
def qa_corpus() -> Dataset:
|
def qa_corpus() -> Dataset:
|
||||||
|
|
@ -18,8 +16,3 @@ def qa_corpus() -> Dataset:
|
||||||
corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories")
|
corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories")
|
||||||
corpus.save_to_disk(ds_path)
|
corpus.save_to_disk(ds_path)
|
||||||
return corpus
|
return corpus
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
|
||||||
def llm_judge() -> LLMJudge:
|
|
||||||
return LLMJudge()
|
|
||||||
|
|
|
||||||
|
|
@ -6,7 +6,7 @@ from llm_judge import LLMJudge
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.qa.ollama import QA
|
from haiku.rag.qa.ollama import QuestionAnswerOllamaAgent
|
||||||
|
|
||||||
db_path = Path(__file__).parent / "data" / "benchmark.sqlite"
|
db_path = Path(__file__).parent / "data" / "benchmark.sqlite"
|
||||||
|
|
||||||
|
|
@ -88,7 +88,7 @@ async def run_qa_benchmark(k: int | None = None):
|
||||||
total_questions = 0
|
total_questions = 0
|
||||||
|
|
||||||
async with HaikuRAG(db_path) as rag:
|
async with HaikuRAG(db_path) as rag:
|
||||||
qa = QA(rag)
|
qa = QuestionAnswerOllamaAgent(rag)
|
||||||
|
|
||||||
for i, doc in enumerate(tqdm(corpus, desc="QA Benchmarking")):
|
for i, doc in enumerate(tqdm(corpus, desc="QA Benchmarking")):
|
||||||
question = doc["question"] # type: ignore
|
question = doc["question"] # type: ignore
|
||||||
|
|
|
||||||
|
|
@ -1,29 +1,52 @@
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from datasets import Dataset
|
from datasets import Dataset
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.qa.ollama import QA
|
from haiku.rag.qa.ollama import QuestionAnswerOllamaAgent
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
try:
|
||||||
import sys
|
from haiku.rag.qa.openai import QuestionAnswerOpenAIAgent
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
sys.path.append(str(Path(__file__).parent))
|
OPENAI_AVAILABLE = True
|
||||||
from llm_judge import LLMJudge
|
except ImportError:
|
||||||
|
QuestionAnswerOpenAIAgent = None
|
||||||
|
OPENAI_AVAILABLE = False
|
||||||
|
|
||||||
|
from .llm_judge import LLMJudge
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_qa_with_dataset_question(qa_corpus: Dataset, llm_judge: "LLMJudge"):
|
async def test_qa_ollama_with_dataset_question(qa_corpus: Dataset):
|
||||||
"""Test QA with actual question from the dataset using LLM judge."""
|
"""Test QA with actual question from the dataset using LLM judge."""
|
||||||
client = HaikuRAG(":memory:")
|
client = HaikuRAG(":memory:")
|
||||||
qa = QA(client)
|
qa = QuestionAnswerOllamaAgent(client)
|
||||||
|
llm_judge = LLMJudge()
|
||||||
|
|
||||||
|
doc = qa_corpus[1]
|
||||||
|
await client.create_document(
|
||||||
|
content=doc["document_extracted"], uri=doc["document_id"]
|
||||||
|
)
|
||||||
|
|
||||||
|
question = doc["question"]
|
||||||
|
expected_answer = doc["answer"]
|
||||||
|
|
||||||
|
answer = await qa.answer(question)
|
||||||
|
is_equivalent = await llm_judge.judge_answers(question, answer, expected_answer)
|
||||||
|
|
||||||
|
assert is_equivalent, (
|
||||||
|
f"Generated answer not equivalent to expected answer.\nQuestion: {question}\nGenerated: {answer}\nExpected: {expected_answer}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI not available")
|
||||||
|
async def test_qa_openai_basic(qa_corpus: Dataset):
|
||||||
|
"""Test OpenAI QA basic functionality."""
|
||||||
|
client = HaikuRAG(":memory:")
|
||||||
|
qa = QuestionAnswerOpenAIAgent(client) # type: ignore
|
||||||
|
llm_judge = LLMJudge()
|
||||||
|
|
||||||
# Use the first document from the corpus
|
|
||||||
doc = qa_corpus[1]
|
doc = qa_corpus[1]
|
||||||
|
|
||||||
# Add the document to database
|
|
||||||
await client.create_document(
|
await client.create_document(
|
||||||
content=doc["document_extracted"], uri=doc["document_id"]
|
content=doc["document_extracted"], uri=doc["document_id"]
|
||||||
)
|
)
|
||||||
|
|
@ -32,11 +55,8 @@ async def test_qa_with_dataset_question(qa_corpus: Dataset, llm_judge: "LLMJudge
|
||||||
expected_answer = doc["answer"]
|
expected_answer = doc["answer"]
|
||||||
|
|
||||||
answer = await qa.answer(question)
|
answer = await qa.answer(question)
|
||||||
# Use LLM judge to evaluate answer equivalence
|
|
||||||
is_equivalent = await llm_judge.judge_answers(question, answer, expected_answer)
|
is_equivalent = await llm_judge.judge_answers(question, answer, expected_answer)
|
||||||
|
|
||||||
assert isinstance(answer, str)
|
|
||||||
assert len(answer) > 0
|
|
||||||
assert is_equivalent, (
|
assert is_equivalent, (
|
||||||
f"Generated answer not equivalent to expected answer.\nQuestion: {question}\nGenerated: {answer}\nExpected: {expected_answer}"
|
f"Generated answer not equivalent to expected answer.\nQuestion: {question}\nGenerated: {answer}\nExpected: {expected_answer}"
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue