bigscience/promptsource
105
1#2# Code for managing session state, which is needed for multi-input forms3# See https://github.com/streamlit/streamlit/issues/15574#5# This code is taken from6# https://gist.github.com/okld/0aba4869ba6fdc8d49132e6974e2e6627#8 9from streamlit.hashing import _CodeHasher10from streamlit.report_thread import get_report_ctx11from streamlit.server.server import Server12 13 14class _SessionState:15 def __init__(self, session, hash_funcs):16 """Initialize SessionState instance."""17 self.__dict__["_state"] = {18 "data": {},19 "hash": None,20 "hasher": _CodeHasher(hash_funcs),21 "is_rerun": False,22 "session": session,23 }24 25 def __call__(self, **kwargs):26 """Initialize state data once."""27 for item, value in kwargs.items():28 if item not in self._state["data"]:29 self._state["data"][item] = value30 31 def __getitem__(self, item):32 """Return a saved state value, None if item is undefined."""33 return self._state["data"].get(item, None)34 35 def __getattr__(self, item):36 """Return a saved state value, None if item is undefined."""37 return self._state["data"].get(item, None)38 39 def __setitem__(self, item, value):40 """Set state value."""41 self._state["data"][item] = value42 43 def __setattr__(self, item, value):44 """Set state value."""45 self._state["data"][item] = value46 47 def clear(self):48 """Clear session state and request a rerun."""49 self._state["data"].clear()50 self._state["session"].request_rerun(None)51 52 def sync(self):53 """54 Rerun the app with all state values up to date from the beginning to55 fix rollbacks.56 """57 data_to_bytes = self._state["hasher"].to_bytes(self._state["data"], None)58 59 # Ensure to rerun only once to avoid infinite loops60 # caused by a constantly changing state value at each run.61 #62 # Example: state.value += 163 if self._state["is_rerun"]:64 self._state["is_rerun"] = False65 66 elif self._state["hash"] is not None:67 if self._state["hash"] != data_to_bytes:68 self._state["is_rerun"] = True69 self._state["session"].request_rerun(None)70 71 self._state["hash"] = data_to_bytes72 73 74def _get_session():75 session_id = get_report_ctx().session_id76 session_info = Server.get_current()._get_session_info(session_id)77 78 if session_info is None:79 raise RuntimeError("Couldn't get your Streamlit Session object.")80 81 return session_info.session82 83 84def _get_state(hash_funcs=None):85 session = _get_session()86 87 if not hasattr(session, "_custom_session_state"):88 session._custom_session_state = _SessionState(session, hash_funcs)89 90 return session._custom_session_state91 