Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions test/test_knowledge.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from typeagent.knowpro import convknowledge, kplib

from fixtures import really_needs_auth
from unittest.mock import AsyncMock


class MockKnowledgeExtractor:
Expand Down Expand Up @@ -45,16 +46,27 @@ async def test_extract_knowledge_from_text(
mock_knowledge_extractor: convknowledge.KnowledgeExtractor,
):
"""Test extracting knowledge from a single text input."""

mock_knowledge_extractor.extract = AsyncMock(
side_effect=mock_knowledge_extractor.extract
)

result = await extract_knowledge_from_text(mock_knowledge_extractor, "test text", 3)
assert isinstance(result, Success)
assert result.value.topics[0] == "test text"

assert mock_knowledge_extractor.extract.call_count == 1

failure_result = await extract_knowledge_from_text(
mock_knowledge_extractor, "error", 3
)
assert isinstance(failure_result, Failure)
assert failure_result.message == "Extraction failed"

assert (
mock_knowledge_extractor.extract.call_count == 5
) # 1 for success | 1 + 3 retries for failures


@pytest.mark.asyncio
async def test_extract_knowledge_from_text_batch(
Expand Down
15 changes: 13 additions & 2 deletions typeagent/knowpro/knowledge.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from . import convknowledge
from . import kplib
from .interfaces import IKnowledgeExtractor
from typechat import Result, Failure, Success


def create_knowledge_extractor(
Expand All @@ -25,8 +26,18 @@ async def extract_knowledge_from_text(
max_retries: int,
) -> Result[kplib.KnowledgeResponse]:
"""Extract knowledge from a single text input with retries."""
# TODO: Add a retry mechanism to handle transient errors.
return await knowledge_extractor.extract(text)

if max_retries < 0:
max_retries = 0

attempt = 0
while attempt <= max_retries:
result = await knowledge_extractor.extract(text)
if isinstance(result, Success) or attempt == max_retries:
return result
attempt += 1

return Failure("No result returned after retries")


async def extract_knowledge_from_text_batch(
Expand Down