diff --git a/burr/core/state.py b/burr/core/state.py index a9a31ac1d..d82e6f4b7 100644 --- a/burr/core/state.py +++ b/burr/core/state.py @@ -22,7 +22,18 @@ import inspect import logging from functools import cached_property -from typing import Any, Callable, Dict, Generic, Iterator, Mapping, Optional, TypeVar, Union +from typing import ( + Any, + Callable, + Dict, + Generic, + Iterator, + List, + Mapping, + Optional, + TypeVar, + Union, +) from burr.core import serde from burr.core.typing import DictBasedTypingSystem, TypingSystem @@ -460,6 +471,14 @@ def __len__(self) -> int: def __iter__(self) -> Iterator[Any]: return iter(self._state) + def keys(self): + """Returns a list of the state keys only (without values for cleaner display). + + Returns: + list: A list of state keys + """ + return list(self._state) + def __repr__(self): return self.get_all().__repr__() # quick hack diff --git a/tests/core/test_state.py b/tests/core/test_state.py index dede0e99e..1c19b8571 100644 --- a/tests/core/test_state.py +++ b/tests/core/test_state.py @@ -226,3 +226,17 @@ def test_state_apply_keeps_typing_system(): state = State({"foo": "bar"}, typing_system=SimpleTypingSystem()) assert state.update(foo="baz").typing_system is state.typing_system assert state.subset("foo").typing_system is state.typing_system + + +def test_state_keys_returns_list(): + """Test that State.keys() returns a list (fixes #409)""" + state = State({"a": 1, "b": 2, "c": 3}) + keys = state.keys() + + # Should return a list with the correct keys + assert isinstance(keys, list) + assert keys == ["a", "b", "c"] + + # Test with empty state + empty_state = State() + assert empty_state.keys() == []