import gradio as gr import time from functools import partial # 从各个模块导入所需的功能 import config from analysis_functions import ( display_transfer_function, time_domain_analysis, frequency_domain_analysis, root_locus_analysis ) # ===== 新增:算例演示模块函数导入 ===== from case_demo_functions import run_case_demo from chatbot import chat_with_ai from user_stats import get_online_status_html, update_user_activity from ui_components import ( create_header, create_time_domain_tab, create_frequency_domain_tab, create_root_locus_tab, create_case_demo_tab, create_chatbot_tab ) # 加载外部CSS文件 with open("assets/styles.css", "r", encoding="utf-8") as f: custom_css = f.read() # --- 主应用界面 --- with gr.Blocks(title="自动控制理论学习网站 - AI+数智平台", css=custom_css) as demo: # 1. 创建UI组件 # 用户会话ID(隐藏组件) session_id = gr.State(value=lambda: str(time.time()) + "_" + str(hash(time.time()))) # 创建头部信息和在线计数器 online_counter = create_header() # 创建共享的输入组件 with gr.Row(): with gr.Column(scale=1): with gr.Group(): gr.HTML("
📊 通用系统参数
") num_input = gr.Textbox( label="传递函数分子系数 (Numerator)", value="1", placeholder="例如: 1 或 1,2,3", info="💡 用逗号分隔多个系数,从最高次项到常数项" ) den_input = gr.Textbox( label="传递函数分母系数 (Denominator)", value="1,6,11,6", placeholder="例如: 1,2,1", info="💡 分母阶数通常高于或等于分子阶数" ) # 创建功能选项卡 with gr.Tabs() as tabs: with gr.TabItem("⏱️ 时域分析 (Time Domain)", id=0): time_domain_ui = create_time_domain_tab() with gr.TabItem("📊 频域分析 (Frequency Domain)", id=1): freq_domain_ui = create_frequency_domain_tab() with gr.TabItem("🎯 根轨迹 (Root Locus)", id=2): root_locus_ui = create_root_locus_tab() # ===== 新增:算例演示 Tab(位于根轨迹与智能问答之间) ===== with gr.TabItem("🧪 算例演示 (Case Demo)", id=3): case_demo_ui = create_case_demo_tab() with gr.TabItem("🤖 智能问答 (Q&A)", id=4): chatbot_ui = create_chatbot_tab() # 2. 绑定事件逻辑 # --- 通用函数 --- # 每次操作前更新用户活跃状态 def wrap_with_activity_update(fn, sid): update_user_activity(sid) # 使用 partial 将 session_id 绑定到函数上 # 这样Gradio调用时就不需要显式传递session_id了 return partial(fn, session_id=sid) # --- 时域分析事件 --- time_domain_ui["confirm_button"].click( fn=display_transfer_function, inputs=[num_input, den_input], outputs=[time_domain_ui["tf_display"]] ).then(lambda: get_online_status_html(), outputs=online_counter) time_domain_ui["analyze_button"].click( fn=time_domain_analysis, inputs=[num_input, den_input], outputs=[time_domain_ui["output_plot"], time_domain_ui["output_metrics"]] ).then(lambda: get_online_status_html(), outputs=online_counter) # --- 频域分析事件 --- def update_frequency_analysis_wrapper(num, den, log_k): k = 10**log_k fig, metrics, tf_latex, stability = frequency_domain_analysis(num, den, k) return fig, metrics, tf_latex, stability, k, get_online_status_html() freq_inputs = [num_input, den_input, freq_domain_ui["log_k_slider"]] freq_outputs = [ freq_domain_ui["plot_output"], freq_domain_ui["metrics_display"], freq_domain_ui["tf_display"], freq_domain_ui["stability_display"], freq_domain_ui["k_number_display"], online_counter ] freq_domain_ui["log_k_slider"].release( fn=update_frequency_analysis_wrapper, inputs=freq_inputs, outputs=freq_outputs ) # --- 根轨迹分析事件 --- def update_rl_view_wrapper(log_k, num, den): fig, poles, k_val = root_locus_analysis(num, den, log_k) return fig, poles, k_val, get_online_status_html() rl_inputs = [root_locus_ui["log_k_slider"], num_input, den_input] rl_outputs = [ root_locus_ui["plot_output"], root_locus_ui["poles_display"], root_locus_ui["k_number_display"], online_counter ] root_locus_ui["log_k_slider"].release( fn=update_rl_view_wrapper, inputs=rl_inputs, outputs=rl_outputs ) # 当输入框变化时,也更新频域和根轨迹(如果它们是当前可见的) def update_all_on_tf_change(num, den, log_k_freq, log_k_rl): # 更新频域 k_freq = 10**log_k_freq fig_freq, metrics, tf_latex, stability = frequency_domain_analysis(num, den, k_freq) # 更新根轨迹 fig_rl, poles, k_val_rl = root_locus_analysis(num, den, log_k_rl) return ( fig_freq, metrics, tf_latex, stability, k_freq, fig_rl, poles, k_val_rl, get_online_status_html() ) tf_change_inputs = [num_input, den_input, freq_domain_ui["log_k_slider"], root_locus_ui["log_k_slider"]] tf_change_outputs = freq_outputs[:-1] + rl_outputs[:-1] + [online_counter] num_input.change(fn=update_all_on_tf_change, inputs=tf_change_inputs, outputs=tf_change_outputs) den_input.change(fn=update_all_on_tf_change, inputs=tf_change_inputs, outputs=tf_change_outputs) # ===== 新增:算例演示事件包装器 ===== def run_case_demo_wrapper(sim_time, dt, initial_soc, initial_engine_power, profile, rpm_scale, load_scale, sid): update_user_activity(sid) fig, summary, table_data = run_case_demo( sim_time_s=sim_time, dt=dt, initial_soc_pct=initial_soc, initial_engine_power_kw=initial_engine_power, profile_name=profile, rpm_scale=rpm_scale, load_scale=load_scale ) return fig, summary, table_data, get_online_status_html() # ===== 新增:算例演示按钮事件绑定 ===== case_demo_ui["run_button"].click( fn=run_case_demo_wrapper, inputs=[ case_demo_ui["sim_time"], case_demo_ui["dt"], case_demo_ui["initial_soc"], case_demo_ui["initial_engine_power"], case_demo_ui["profile"], case_demo_ui["rpm_scale"], case_demo_ui["load_scale"], session_id ], outputs=[ case_demo_ui["plot"], case_demo_ui["summary"], case_demo_ui["table"], online_counter ] ) # --- 聊天机器人事件 --- async def chat_wrapper(message, history, sid): update_user_activity(sid) # chat_with_ai 是一个生成器,Gradio可以直接处理 async for response in chat_with_ai(message, history): yield response chatbot_ui["send_button"].click( fn=chat_wrapper, inputs=[chatbot_ui["chat_input"], chatbot_ui["chatbot"], session_id], outputs=chatbot_ui["chatbot"] ).then(lambda: ("", get_online_status_html()), outputs=[chatbot_ui["chat_input"], online_counter]) chatbot_ui["chat_input"].submit( fn=chat_wrapper, inputs=[chatbot_ui["chat_input"], chatbot_ui["chatbot"], session_id], outputs=chatbot_ui["chatbot"] ).then(lambda: ("", get_online_status_html()), outputs=[chatbot_ui["chat_input"], online_counter]) def clear_chat_wrapper(sid): update_user_activity(sid) return [], get_online_status_html() chatbot_ui["clear_button"].click( fn=clear_chat_wrapper, inputs=[session_id], outputs=[chatbot_ui["chatbot"], online_counter] ) # --- 页面加载和定时器事件 --- def on_page_load(sid): update_user_activity(sid) return get_online_status_html() demo.load(fn=on_page_load, inputs=[session_id], outputs=[online_counter]) gr.Timer(10).tick(fn=get_online_status_html, outputs=online_counter) if __name__ == "__main__": demo.queue().launch( server_name=config.SERVER_NAME, server_port=config.SERVER_PORT, share=config.SHARE )