참고소스 수정본
This commit is contained in:
@@ -0,0 +1,67 @@
|
||||
import pytest
|
||||
from guardrails import Guard
|
||||
from typing import List, Optional
|
||||
from tests.integration_tests.test_assets.validators import RegexMatch
|
||||
|
||||
pytest.importorskip("llama_index")
|
||||
|
||||
from llama_index.core.chat_engine.types import ( # noqa
|
||||
BaseChatEngine, # noqa
|
||||
AgentChatResponse, # noqa
|
||||
StreamingAgentChatResponse, # noqa
|
||||
) # noqa
|
||||
from llama_index.core.base.llms.types import ChatMessage # noqa
|
||||
from guardrails.integrations.llama_index import GuardrailsChatEngine # noqa
|
||||
|
||||
|
||||
class MockChatEngine(BaseChatEngine):
|
||||
def chat(
|
||||
self, message: str, chat_history: Optional[List[ChatMessage]] = None
|
||||
) -> AgentChatResponse:
|
||||
return AgentChatResponse(response="Mock response")
|
||||
|
||||
async def achat(
|
||||
self, message: str, chat_history: Optional[List[ChatMessage]] = None
|
||||
) -> AgentChatResponse:
|
||||
return AgentChatResponse(response="Mock async chat response")
|
||||
|
||||
def stream_chat(
|
||||
self, message: str, chat_history: Optional[List[ChatMessage]] = None
|
||||
):
|
||||
return StreamingAgentChatResponse(response="Mock stream chat response")
|
||||
|
||||
async def astream_chat(
|
||||
self, message: str, chat_history: Optional[List[ChatMessage]] = None
|
||||
):
|
||||
return StreamingAgentChatResponse(response="Mock async stream chat response")
|
||||
|
||||
@property
|
||||
def chat_history(self) -> List[ChatMessage]:
|
||||
return []
|
||||
|
||||
def reset(self):
|
||||
pass
|
||||
|
||||
|
||||
pytest.importorskip("llama_index")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def guard():
|
||||
return Guard().use(RegexMatch("Mock response", match_type="search"))
|
||||
|
||||
|
||||
class TestGuardrailsChatEngine:
|
||||
def test_guardrails_engine_init(self, guard):
|
||||
engine = MockChatEngine()
|
||||
guardrails_engine = GuardrailsChatEngine(engine, guard)
|
||||
assert isinstance(guardrails_engine, GuardrailsChatEngine)
|
||||
assert guardrails_engine.guard == guard
|
||||
|
||||
def test_guardrails_engine_chat(self, guard):
|
||||
engine = MockChatEngine()
|
||||
guardrails_engine = GuardrailsChatEngine(engine, guard)
|
||||
|
||||
result = guardrails_engine.chat("Mock response")
|
||||
assert isinstance(result, AgentChatResponse)
|
||||
assert result.response == "Mock response"
|
||||
@@ -0,0 +1,63 @@
|
||||
import pytest
|
||||
from guardrails import Guard
|
||||
from guardrails.errors import ValidationError
|
||||
from typing import Optional
|
||||
from tests.integration_tests.test_assets.validators import RegexMatch
|
||||
|
||||
pytest.importorskip("llama_index")
|
||||
|
||||
from llama_index.core.query_engine import BaseQueryEngine # noqa
|
||||
from llama_index.core.schema import QueryBundle # noqa
|
||||
from llama_index.core.base.response.schema import Response # noqa
|
||||
from llama_index.core.prompts.mixin import PromptMixinType # noqa
|
||||
from llama_index.core.callbacks import CallbackManager # noqa
|
||||
|
||||
|
||||
class MockQueryEngine(BaseQueryEngine):
|
||||
def __init__(self, callback_manager: Optional[CallbackManager] = None):
|
||||
super().__init__(callback_manager)
|
||||
|
||||
def _query(self, query_bundle: QueryBundle) -> Response:
|
||||
return Response(response="Mock response")
|
||||
|
||||
async def _aquery(self, query_bundle: QueryBundle) -> Response:
|
||||
return Response(response="Mock async query response")
|
||||
|
||||
def _get_prompt_modules(self) -> PromptMixinType:
|
||||
return {}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def guard():
|
||||
return Guard().use(RegexMatch("Mock response", match_type="search"))
|
||||
|
||||
|
||||
class TestGuardrailsQueryEngine:
|
||||
def test_guardrails_engine_init(self, guard):
|
||||
from guardrails.integrations.llama_index import GuardrailsQueryEngine
|
||||
|
||||
engine = MockQueryEngine()
|
||||
guardrails_engine = GuardrailsQueryEngine(engine, guard)
|
||||
assert isinstance(guardrails_engine, GuardrailsQueryEngine)
|
||||
assert guardrails_engine.guard == guard
|
||||
|
||||
def test_guardrails_engine_query(self, guard):
|
||||
from guardrails.integrations.llama_index import GuardrailsQueryEngine
|
||||
|
||||
engine = MockQueryEngine()
|
||||
guardrails_engine = GuardrailsQueryEngine(engine, guard)
|
||||
|
||||
result = guardrails_engine._query(QueryBundle(query_str="Mock response"))
|
||||
assert isinstance(result, Response)
|
||||
assert result.response == "Mock response"
|
||||
|
||||
def test_guardrails_engine_query_validation_failure(self, guard):
|
||||
from guardrails.integrations.llama_index import GuardrailsQueryEngine
|
||||
|
||||
engine = MockQueryEngine()
|
||||
guardrails_engine = GuardrailsQueryEngine(engine, guard)
|
||||
|
||||
engine._query = lambda _: Response(response="Invalid response")
|
||||
|
||||
with pytest.raises(ValidationError, match="Validation failed"):
|
||||
guardrails_engine._query(QueryBundle(query_str="Invalid query"))
|
||||
Reference in New Issue
Block a user