diff --git a/src/rai_core/rai/agents/langchain/callback.py b/src/rai_core/rai/agents/langchain/callback.py index 4978598ce..dc9c13fa0 100644 --- a/src/rai_core/rai/agents/langchain/callback.py +++ b/src/rai_core/rai/agents/langchain/callback.py @@ -39,6 +39,10 @@ def __init__( self.stream_response = stream_response self.splitting_chars = splitting_chars or ["\n", ".", "!", "?"] self.chunks_buffer = "" + if not isinstance(max_buffer_size, int) or isinstance(max_buffer_size, bool): + raise ValueError("max_buffer_size must be a positive int") + if max_buffer_size <= 0: + raise ValueError("max_buffer_size must be a positive int") self.max_buffer_size = max_buffer_size self._buffer_lock = threading.Lock() self.logger = logger or logging.getLogger(__name__) diff --git a/tests/agents/langchain/test_hri_callback_max_buffer.py b/tests/agents/langchain/test_hri_callback_max_buffer.py new file mode 100644 index 000000000..45b4ee7f4 --- /dev/null +++ b/tests/agents/langchain/test_hri_callback_max_buffer.py @@ -0,0 +1,18 @@ +# Copyright (C) 2026 Robotec.AI +import pytest + +from rai.agents.langchain.callback import HRICallbackHandler + + +def test_hri_callback_rejects_non_positive_max_buffer_size(): + with pytest.raises(ValueError, match="max_buffer_size"): + HRICallbackHandler(connectors={}, max_buffer_size=0) + with pytest.raises(ValueError, match="max_buffer_size"): + HRICallbackHandler(connectors={}, max_buffer_size=-5) + with pytest.raises(ValueError, match="max_buffer_size"): + HRICallbackHandler(connectors={}, max_buffer_size=True) # type: ignore[arg-type] + + +def test_hri_callback_accepts_positive_max_buffer_size(): + h = HRICallbackHandler(connectors={}, max_buffer_size=16) + assert h.max_buffer_size == 16