-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
249 lines (216 loc) · 10.4 KB
/
Copy pathmain.py
File metadata and controls
249 lines (216 loc) · 10.4 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
import asyncio
import argparse
import os
import sys
import faulthandler
from pathlib import Path
import warnings
warnings.filterwarnings(
"ignore",
message='Field name ".*" shadows an attribute in parent "Operation"',
category=UserWarning,
module='pydantic',
)
faulthandler.enable()
from dotenv import load_dotenv
# Load .env safely. Some Windows editors save .env as UTF-16, which crashes python-dotenv.
try:
load_dotenv(encoding="utf-8")
except UnicodeDecodeError:
print("WARNING: .env is not UTF-8. Ignoring .env for local-first boot. Rename or resave it as UTF-8.")
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
from src.utils.logger import get_logger
from src.core.config import BrainConfig
from src.core.brain import AIVtuberBrain
from src.modules.llm.gemini_llm import GeminiLLM
from src.modules.llm.glm_llm import GLM47LLM
from src.modules.llm.openai_llm import OpenAILLM
from src.modules.llm.groq_llm import GroqLLM
from src.modules.llm.ollama_provider import OllamaLLM
from src.modules.obs.obs_websocket import OBSController
logger = get_logger("bea")
def parse_args():
parser = argparse.ArgumentParser(description="ProjectBEA - AI Vtuber Engine")
parser.add_argument("--web", action="store_true", help="Start Web Interface (FastAPI + React)")
parser.add_argument("--system-file", default=None, help="Path to system prompt file")
parser.add_argument("--png-dir", default=None, help="Directory for avatar PNGs")
parser.add_argument("--llm-provider", choices=["gemini", "glm", "openai", "groq", "ollama"], default=None, help="LLM Provider to use")
parser.add_argument("--ollama-model", default=None, help="Ollama model")
parser.add_argument("--ollama-timeout", type=float, default=None, help="Ollama timeout seconds")
parser.add_argument("--ollama-host", default=None, help="Ollama host URL")
parser.add_argument("--gemini-key", default=None, help="Google GenAI API Key")
parser.add_argument("--gemini-model", default=None, help="Gemini Model")
parser.add_argument("--glm-key", default=None, help="GLM API Key")
parser.add_argument("--glm-model", default=None, help="GLM Model")
parser.add_argument("--openai-key", default=None, help="OpenAI API Key")
parser.add_argument("--openai-model", default=None, help="OpenAI Model")
parser.add_argument("--groq-key", default=None, help="Groq API Key")
parser.add_argument("--groq-model", default=None, help="Groq Model")
parser.add_argument("--stt-provider", choices=["groq", "none"], default=None, help="STT Provider")
parser.add_argument("--stt-model", default=None, help="STT Model")
parser.add_argument("--obs-host", default=None, help="OBS WebSocket host")
parser.add_argument("--obs-port", type=int, default=None, help="OBS WebSocket port")
parser.add_argument("--obs-password", default=None, help="OBS WebSocket password")
parser.add_argument("--obs-avatar-source", default=None, required=False, help="OBS Source Name for Avatar")
parser.add_argument("--obs-source-type", choices=["image", "media"], default=None, help="OBS Source Type")
parser.add_argument("--obs-text-source", default=None, help="OBS Source Name for Text Bubble")
parser.add_argument("--tts-provider", choices=["edge", "coqui", "orpheus", "kokoro"], default=None, help="TTS Provider")
parser.add_argument("--tts-voice", default=None, help="EdgeTTS Voice")
parser.add_argument("--orpheus-key", default=None, help="Orpheus API Key")
parser.add_argument("--orpheus-endpoint", default=None, help="Orpheus Endpoint")
parser.add_argument("--orpheus-voice", default=None, help="Orpheus Voice")
parser.add_argument("--kokoro-file", default=None, help="Kokoro Model File")
parser.add_argument("--kokoro-voices", default=None, help="Kokoro Voices File")
parser.add_argument("--device-id", type=int, default=None, help="Audio Output Device ID")
parser.add_argument("--typing-delay", type=float, default=None, help="Typing animation delay")
return parser.parse_args()
def _apply_env_secret_fallbacks(config: BrainConfig) -> None:
"""Fill secret config fields from environment without writing them to config.json."""
env_map = {
"gemini_key": "GEMINI_API_KEY",
"glm_key": "GLM_API_KEY",
"openai_key": "OPENAI_API_KEY",
"groq_key": "GROQ_API_KEY",
"orpheus_key": "ORPHEUS_API_KEY",
"orpheus_endpoint": "ORPHEUS_ENDPOINT",
}
for field_name, env_name in env_map.items():
current = getattr(config, field_name, None)
if current:
continue
value = os.getenv(env_name)
if value:
setattr(config, field_name, value)
def _initialize_stt(config: BrainConfig):
"""Initialize STT safely. Returns None when disabled or unavailable."""
provider = (getattr(config, "stt_provider", "none") or "none").strip().lower()
config.stt_provider = provider
if provider != "groq":
logger.info(f"STT disabled/provider={provider}")
return None
if not getattr(config, "groq_key", None):
config.groq_key = os.getenv("GROQ_API_KEY")
if not config.groq_key:
logger.warning("GROQ_API_KEY missing. STT disabled for local-first boot.")
logger.info("STT Provider=groq GroqKeyPresent=False STTLoaded=False")
return None
try:
from src.modules.STT.groq_stt import GroqSTT
except ModuleNotFoundError:
# Some older trees used src/STT instead of src/modules/STT.
from src.STT.groq_stt import GroqSTT
logger.info("Initializing Groq STT...")
stt = GroqSTT(config)
logger.info(f"STT Provider=groq GroqKeyPresent=True STTLoaded={stt is not None}")
return stt
async def main():
args = parse_args()
config = BrainConfig()
cli_overrides = {
"system_prompt_path": args.system_file,
"png_dir": args.png_dir,
"llm_provider": args.llm_provider,
"ollama_model": args.ollama_model,
"ollama_timeout": args.ollama_timeout,
"ollama_host": args.ollama_host,
"gemini_key": args.gemini_key,
"gemini_model": args.gemini_model,
"glm_key": args.glm_key,
"glm_model": args.glm_model,
"openai_key": args.openai_key,
"openai_model": args.openai_model,
"groq_key": args.groq_key,
"groq_model": args.groq_model,
"stt_provider": args.stt_provider,
"stt_model": args.stt_model,
"obs_host": args.obs_host,
"obs_port": args.obs_port,
"obs_password": args.obs_password,
"obs_avatar_source": args.obs_avatar_source,
"obs_source_type": args.obs_source_type,
"obs_text_source": args.obs_text_source,
"tts_provider": args.tts_provider,
"tts_voice": args.tts_voice,
"orpheus_key": args.orpheus_key,
"orpheus_endpoint": args.orpheus_endpoint,
"orpheus_voice": args.orpheus_voice,
"kokoro_model": args.kokoro_file,
"kokoro_voices_file": args.kokoro_voices,
"audio_device_id": args.device_id,
"typing_delay": args.typing_delay,
}
for field_name, value in cli_overrides.items():
if value is not None:
setattr(config, field_name, value)
if field_name.endswith("_key"):
logger.info(f"CLI override: {field_name} = [set]")
else:
logger.info(f"CLI override: {field_name} = {value}")
_apply_env_secret_fallbacks(config)
stt = _initialize_stt(config)
if config.llm_provider == "ollama":
llm = OllamaLLM(
model_name=config.ollama_model,
timeout=config.ollama_timeout,
host=config.ollama_host,
stt_interface=stt,
)
elif config.llm_provider == "gemini":
if not config.gemini_key:
logger.error("GEMINI_API_KEY is missing via env, config, or CLI.")
return
llm = GeminiLLM(api_key=config.gemini_key, model_name=config.gemini_model)
elif config.llm_provider == "glm":
if not config.glm_key:
logger.error("GLM_API_KEY is missing via env, config, or CLI.")
return
llm = GLM47LLM(api_key=config.glm_key, model_name=config.glm_model, stt_interface=stt)
elif config.llm_provider == "openai":
if not config.openai_key:
logger.error("OPENAI_API_KEY is missing via env, config, or CLI.")
return
llm = OpenAILLM(api_key=config.openai_key, model_name=config.openai_model, stt_interface=stt)
elif config.llm_provider == "groq":
if not config.groq_key:
logger.error("GROQ_API_KEY is missing via env, config, or CLI.")
return
llm = GroqLLM(api_key=config.groq_key, model_name=config.groq_model, stt_interface=stt)
else:
logger.error(f"Unknown LLM provider: {config.llm_provider}")
return
if config.tts_provider == "orpheus":
from src.modules.tts.orpheus_tts_wrapper import OrpheusTTSWrapper
tts = OrpheusTTSWrapper(api_key=config.orpheus_key, endpoint_url=config.orpheus_endpoint, voice=config.orpheus_voice)
elif config.tts_provider == "kokoro":
from src.modules.tts.kokoro_tts_wrapper import KokoroTTSWrapper
tts = KokoroTTSWrapper(
model_path=config.kokoro_model,
voices_path=config.kokoro_voices_file,
voice=config.kokoro_voice,
speed=config.kokoro_speed,
lang=config.kokoro_lang,
)
else:
from src.modules.tts.edge_tts_wrapper import EdgeTTSWrapper
tts = EdgeTTSWrapper(voice=config.tts_voice, pitch=config.tts_pitch, rate=config.tts_rate, volume=config.tts_volume)
obs = OBSController(host=config.obs_host, port=config.obs_port, password=config.obs_password, source_name=config.obs_avatar_source)
brain = AIVtuberBrain(config, llm, tts, stt, obs)
try:
brain.initialize()
await brain.start_skills()
if args.web:
from src.web.server import run_server
logger.info("Starting Web Interface at http://localhost:8000")
await run_server(brain, port=8000)
else:
await brain.run_loop()
except KeyboardInterrupt:
logger.info("Stopping...")
if brain.memory_skill and brain.memory_skill.enabled:
logger.info("Saving pending memories...")
await brain.memory_skill.save_all_pending()
finally:
await brain.skill_manager.stop()
brain.shutdown()
if __name__ == "__main__":
asyncio.run(main())