mentalHealthPatientenv / server /mentalHealthPatientenv_environment.py
Code0ut's picture
Upload folder using huggingface_hub
388d0ec verified
Raw
History Blame Contribute Delete
7.67 kB
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
"""
Mentalhealthpatientenv Environment Implementation.
A simple test environment that echoes back messages sent to it.
Perfect for testing HTTP server infrastructure.
"""
import random
from uuid import uuid4
from openenv.core.env_server.interfaces import Environment
try:
from server.patient_mental_state import MentalState
except ImportError:
from patient_mental_state import MentalState
from server.disorders import get_random_disorder
from server.reward import calculate_reward
from server.response_generator import generate_response
try:
from ..models import MentalhealthpatientenvAction, MentalhealthpatientenvObservation
except ImportError:
from models import MentalhealthpatientenvAction, MentalhealthpatientenvObservation
class MentalhealthpatientenvEnvironment(Environment):
"""
A simple echo environment that echoes back messages.
This environment is designed for testing the HTTP server infrastructure.
It maintains minimal state and simply echoes back whatever message it receives.
Example:
>>> env = MentalhealthpatientenvEnvironment()
>>> obs = env.reset()
>>> print(obs.echoed_message) # "Mentalhealthpatientenv environment ready!"
>>>
>>> obs = env.step(MentalhealthpatientenvAction(message="Hello"))
>>> print(obs.echoed_message) # "Hello"
>>> print(obs.message_length) # 5
"""
# Enable concurrent WebSocket sessions.
# Set to True if your environment isolates state between instances.
# When True, multiple WebSocket clients can connect simultaneously, each
# getting their own environment instance (when using factory mode in app.py).
SUPPORTS_CONCURRENT_SESSIONS: bool = True
def __init__(self):
"""Initialize the mentalHealthPatientenv environment."""
self._state = MentalState(episode_id=str(uuid4()), step_count=0)
self._reset_count = 0
def reset(self) -> MentalhealthpatientenvObservation:
"""
Reset the environment.
Returns:
MentalhealthpatientenvObservation with initial message and state information.
"""
self._state = MentalState(episode_id=str(uuid4()), step_count=0)
self._reset_count += 1
return MentalhealthpatientenvObservation(
response="Mentalhealthpatientenv environment ready!",
clarity=random.uniform(0.0, 1.0),
emotional_state="neutral",
trust_level=random.uniform(0.0, 0.5),
risk_flag=False,
done=False,
reward=0.0,
metadata={"reset_count": self._reset_count}
)
def step(self, action: MentalhealthpatientenvAction) -> MentalhealthpatientenvObservation: # type: ignore[override]
"""
Execute a step in the environment by echoing the message.
Args:
action: MentalhealthpatientenvAction containing the message to echo
Returns:
MentalhealthpatientenvObservation with the echoed message and its length
"""
self._state.step_count += 1
if action.action_type == "Initialize":
self._state.action_sequence = []
if action.task_difficulty == "easy":
self._state.disorder = get_random_disorder(1)
self._state.severity = "mild"
self._state.trust_level = random.uniform(0.4, 0.8)
self._state.disclosed_risk = True
elif action.task_difficulty == "medium":
self._state.disorder = get_random_disorder(1)
self._state.severity = "moderate"
self._state.trust_level = random.uniform(0.2, 0.6)
self._state.disclosed_risk = random.choice([True, False])
elif action.task_difficulty == "hard":
self._state.disorder = get_random_disorder(3)
self._state.severity = "severe"
self._state.trust_level = random.uniform(0.0, 0.4)
self._state.disclosed_risk = False
else:
print(f"Unknown task difficulty: {action.task_difficulty}. Using default state.")
return MentalhealthpatientenvObservation(
response="Environment initialized with patient state.",
clarity=random.uniform(0.0, 1.0),
emotional_state="neutral",
trust_level=self._state.trust_level,
risk_flag=self._state.disclosed_risk,
done=False,
reward=0.0,
metadata={"initialized_state": self._state.dict()}
)
elif action.action_type == "diagnose":
# Simulate a diagnosis action by providing feedback based on the patient's state
self._state.action_sequence.append(action.action_type)
curr_diagnosis=action.message.strip().split(',')
actual_disorders = self._state.disorder
reward = calculate_reward(
disorder=actual_disorders,
action_type=action.action_type,
task_difficulty=action.task_difficulty,
diagnosis=curr_diagnosis,
state=self._state,
action_sequence=self._state.action_sequence
)
return MentalhealthpatientenvObservation(
response=generate_response(curr_diagnosis, actual_disorders,state=self._state ) if self._state.step_count < self._state.max_turns else "Maximum turns reached. Ending session.",
clarity=random.uniform(0.0, 1.0),
emotional_state="neutral",
trust_level=self._state.trust_level,
risk_flag=self._state.disclosed_risk,
done=True if reward >= 0.8 or self._state.step_count >= self._state.max_turns else False,
reward=reward,
metadata={"diagnosis_result": curr_diagnosis, "actual_disorders": actual_disorders, "reward": reward}
)
elif action.action_type in ["ask_open", "ask_direct", "ask_risk", "reflect"]:
self._state.action_sequence.append(action.action_type)
reward = calculate_reward(
disorder=self._state.disorder,
action_type=action.action_type,
task_difficulty=action.task_difficulty,
state=self._state,
action_sequence=self._state.action_sequence
)
return MentalhealthpatientenvObservation(
response=generate_response(action.action_type, action.message, state=self._state) if self._state.step_count < self._state.max_turns else "Maximum turns reached. Ending session.",
clarity=random.uniform(0.0, 1.0),
emotional_state="neutral",
trust_level=self._state.trust_level,
risk_flag=self._state.disclosed_risk,
done=True if self._state.step_count >= self._state.max_turns else False,
reward=reward, # small reward for taking an action
metadata={"received_action": action.action_type}
)
@property
def state(self) -> MentalState:
"""
Get the current environment state.
Returns:
Current State with episode_id and step_count
"""
return self._state