openenv/chat_env
0
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7"""8FastAPI application for the Chat Environment.9 10This module creates an HTTP server that exposes the ChatEnvironment11over HTTP and WebSocket endpoints, compatible with EnvClient.12 13Note: This server requires a tokenizer to be initialized. The tokenizer14must be specified when starting the server.15 16Usage:17 # Development (with auto-reload):18 uvicorn envs.chat_env.server.app:app --reload --host 0.0.0.0 --port 800019 20 # Production:21 uvicorn envs.chat_env.server.app:app --host 0.0.0.0 --port 8000 --workers 422 23 # Or run directly:24 python -m envs.chat_env.server.app25"""26 27import os28 29from openenv.core.env_server import create_app30 31# Support both in-repo and standalone imports32try:33 # In-repo imports (when running from OpenEnv repository)34 from ..models import ChatAction, ChatObservation35 from .chat_environment import ChatEnvironment36except ImportError as e:37 if "relative import" not in str(e) and "no known parent package" not in str(e):38 raise39 # Standalone imports (when running via uvicorn server.app:app)40 from models import ChatAction, ChatObservation41 from server.chat_environment import ChatEnvironment42 43 44# Initialize tokenizer based on environment variable45def get_tokenizer():46 """Get tokenizer from environment or use a mock for testing."""47 tokenizer_name = os.environ.get("TOKENIZER_NAME", "gpt2")48 49 try:50 from transformers import AutoTokenizer51 52 tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)53 print(f"Loaded tokenizer: {tokenizer_name}")54 return tokenizer55 except ImportError:56 print(57 "Warning: transformers not installed, using mock tokenizer for testing only"58 )59 # Use mock tokenizer from tests60 import sys61 from pathlib import Path62 63 # Add parent directory to path to import test utilities64 test_path = Path(__file__).parent65 sys.path.insert(0, str(test_path))66 67 from test_chat_env import MockTokenizer68 69 return MockTokenizer()70 71 72# Get system prompt from environment73system_prompt = os.environ.get("SYSTEM_PROMPT", None)74 75 76# Factory function to create ChatEnvironment instances77def create_chat_environment():78 """Factory function that creates ChatEnvironment with tokenizer."""79 tokenizer = get_tokenizer()80 return ChatEnvironment(tokenizer=tokenizer, system_prompt=system_prompt)81 82 83# Create the FastAPI app with web interface and README integration84# Pass the factory function instead of an instance for WebSocket session support85app = create_app(86 create_chat_environment, ChatAction, ChatObservation, env_name="chat_env"87)88 89 90def main():91 import uvicorn92 93 uvicorn.run(app, host="0.0.0.0", port=8000)94 95 96if __name__ == "__main__":97 main()98 