Skip to content
Open
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
2 changes: 2 additions & 0 deletions nemoguardrails/library/hf_classifier/flows.co
Original file line number Diff line number Diff line change
Expand Up @@ -22,4 +22,6 @@ flow hf classifier check retrieval $classifier
$response = await HfClassifierCheckRetrievalAction(classifier=$classifier)

if $response.is_transform
# FIX: Added global declaration to propagate the change to calling context
global $relevant_chunks
$relevant_chunks = $response.transform_text["relevant_chunks"]
4 changes: 3 additions & 1 deletion nemoguardrails/library/hf_classifier/flows.v1.co
Original file line number Diff line number Diff line change
Expand Up @@ -25,4 +25,6 @@ define subflow hf classifier check retrieval
$response = execute hf_classifier_check_retrieval(classifier=$classifier)

if $response.is_transform
$relevant_chunks = $response.transform_text["relevant_chunks"]
# In Colang 1.0, subflows share global scope.
# The caller must initialize $relevant_chunks before calling this subflow.
$relevant_chunks = $response.transform_text["relevant_chunks"]
195 changes: 195 additions & 0 deletions tests/test_hf_classifier_retrieval_global.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,195 @@
import unittest
from unittest.mock import AsyncMock

from nemoguardrails import LLMRails, RailsConfig


class MockHfClassifierResponse:
def __init__(self, is_transform=False, transform_text=None, is_blocked=False):
self.is_transform = is_transform
self.transform_text = transform_text or {}
self.is_blocked = is_blocked


class TestHfClassifierRetrievalGlobal(unittest.TestCase):

def test_v1_retrieval_transform_propagates_global(self):
mock_action = AsyncMock()
mock_action.return_value = MockHfClassifierResponse(
is_transform=True,
transform_text={
"relevant_chunks": ["new_chunk_1", "new_chunk_2"],
"response": "Transformed response"
}
)

config = RailsConfig.from_content(
colang_content="""
define user express greeting
"Hello"

define flow test_retrieval
$user_input = "test query"
$classifier = "test_classifier"
$relevant_chunks = execute hf classifier check retrieval
$test_output = $relevant_chunks
""",
yaml_content="""
models:
- type: main
engine: openai
model: gpt-3.5-turbo
"""
)

app = LLMRails(config)
app.register_action("hf_classifier_check_retrieval", mock_action)
app.generate(messages=[{"role": "user", "content": "Hello"}])

self.assertEqual(app.context.get("relevant_chunks"), ["new_chunk_1", "new_chunk_2"])
self.assertEqual(app.context.get("test_output"), ["new_chunk_1", "new_chunk_2"])

def test_original_retrieval_transform_propagates_global(self):
mock_action = AsyncMock()
mock_action.return_value = MockHfClassifierResponse(
is_transform=True,
transform_text={
"relevant_chunks": ["original_chunk_1", "original_chunk_2"],
"response": "Original transformed response"
}
)

config = RailsConfig.from_content(
colang_content="""
define user express greeting
"Hello"

define flow test_retrieval
$user_input = "test query"
$classifier = "test_classifier"
execute hf classifier check retrieval $classifier
$test_output = $relevant_chunks
""",
yaml_content="""
models:
- type: main
engine: openai
model: gpt-3.5-turbo
"""
)

app = LLMRails(config)
app.register_action("HfClassifierCheckRetrievalAction", mock_action)
app.generate(messages=[{"role": "user", "content": "Hello"}])

self.assertEqual(app.context.get("relevant_chunks"), ["original_chunk_1", "original_chunk_2"])
self.assertEqual(app.context.get("test_output"), ["original_chunk_1", "original_chunk_2"])

def test_no_transform_doesnt_modify_chunks(self):
mock_action = AsyncMock()
mock_action.return_value = MockHfClassifierResponse(
is_transform=False,
is_blocked=False
)

config = RailsConfig.from_content(
colang_content="""
define user express greeting
"Hello"

define flow test_no_transform
$user_input = "test query"
$original_chunks = ["original_chunk_1", "original_chunk_2"]
$relevant_chunks = $original_chunks
$classifier = "test_classifier"
execute hf classifier check retrieval
$test_output = $relevant_chunks
""",
yaml_content="""
models:
- type: main
engine: openai
model: gpt-3.5-turbo
"""
)

app = LLMRails(config)
app.register_action("hf_classifier_check_retrieval", mock_action)
app.generate(messages=[{"role": "user", "content": "Hello"}])

self.assertEqual(app.context.get("relevant_chunks"), ["original_chunk_1", "original_chunk_2"])
self.assertEqual(app.context.get("test_output"), ["original_chunk_1", "original_chunk_2"])

def test_v1_retrieval_with_empty_transform_text(self):
mock_action = AsyncMock()
mock_action.return_value = MockHfClassifierResponse(
is_transform=True,
transform_text={"relevant_chunks": []}
)

config = RailsConfig.from_content(
colang_content="""
define user express greeting
"Hello"

define flow test_empty_transform
$user_input = "test query"
$original_chunks = ["original"]
$relevant_chunks = $original_chunks
$classifier = "test_classifier"
$relevant_chunks = execute hf classifier check retrieval
$test_output = $relevant_chunks
""",
yaml_content="""
models:
- type: main
engine: openai
model: gpt-3.5-turbo
"""
)

app = LLMRails(config)
app.register_action("hf_classifier_check_retrieval", mock_action)
app.generate(messages=[{"role": "user", "content": "Hello"}])

self.assertEqual(app.context.get("relevant_chunks"), [])
self.assertEqual(app.context.get("test_output"), [])

def test_v1_retrieval_blocked_flow(self):
mock_action = AsyncMock()
mock_action.return_value = MockHfClassifierResponse(
is_transform=False,
is_blocked=True
)

config = RailsConfig.from_content(
colang_content="""
define user express greeting
"Hello"

define flow test_blocked
$user_input = "test query"
$original_chunks = ["original"]
$relevant_chunks = $original_chunks
$classifier = "test_classifier"
$config.enable_rails_exceptions = True
execute hf classifier check retrieval
$test_output = $relevant_chunks
""",
yaml_content="""
models:
- type: main
engine: openai
model: gpt-3.5-turbo
"""
)

app = LLMRails(config)
app.register_action("hf_classifier_check_retrieval", mock_action)
app.generate(messages=[{"role": "user", "content": "Hello"}])

self.assertEqual(app.context.get("relevant_chunks"), ["original"])


if __name__ == "__main__":
unittest.main()