| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524 |
- # -*- coding: utf-8 -*-
- import os
- import re
- import glob
- import json
- import shutil
- import base64
- import uuid
- import datetime
- import subprocess
- import html
- import streamlit as st
- import streamlit.components.v1 as components
- from openai import OpenAI
- # ========== 启动文件服务(非阻塞)==========
- try:
- if "fastpdf_spawned" not in st.session_state:
- subprocess.Popen(["python", "fastpdf.py"])
- st.session_state["fastpdf_spawned"] = True
- except Exception:
- pass
- # ========== 示例:模型和相应的 API 密钥 ==========
- # 请在生产环境中改为安全地注入 API Key(如通过环境变量)
- default_key = "sk-re2NlaKIQn11ZNWzAbB6339cEbF94c6aAfC8B7Ab82879bEa" # 默认令牌,请勿用于生产
- model_keys = {
- "grok-3": default_key,
- "grok-4.1": default_key,
- "grok-4.2": default_key,
- "gpt-5.4-mini": default_key,
- "gpt-4o-mini": default_key,
- # "gpt-4.1-mini-2025-04-14": default_key,
- "o1-mini": default_key,
- "o4-mini": default_key,
- "deepseek-v3.1": default_key,
- "deepseek-r1": default_key,
- "grok-4": default_key,
- "gpt-5-all": default_key,
- "gpt-4o-all": default_key,
- # "gpt-5-mini-2025-08-07": default_key,
- "o3-mini-all": default_key,
- # "claude-sonnet-4-20250514": default_key,
- }
- st.title("ChatGPT-like Clone")
- selected_model = st.sidebar.selectbox("选择模型", list(model_keys.keys()))
- api_key = model_keys[selected_model]
- api_url = "https://yunwu.ai/v1"
- client = OpenAI(api_key=api_key, base_url=api_url)
- data_dir = "data"
- backup_dir = "data_bak"
- blog_dir = "blog"
- upload_dir = "uploads"
- for directory in [data_dir, backup_dir, blog_dir, upload_dir]:
- if not os.path.exists(directory):
- os.makedirs(directory)
- session_id_file = os.path.join(data_dir, "session_id.txt")
- def load_session_id():
- if os.path.exists(session_id_file):
- try:
- with open(session_id_file, "r", encoding="utf-8") as f:
- c = f.read().strip()
- return int(c) if c.isdigit() else 0
- except Exception:
- return 0
- return 0
- def save_session_id(session_id: int):
- with open(session_id_file, "w", encoding="utf-8") as f:
- f.write(str(session_id))
- if "openai_model" not in st.session_state:
- st.session_state["openai_model"] = selected_model
- def load_history(file_path: str):
- match = re.search(r'chat_history_(\d+)\.json', os.path.basename(file_path))
- if match:
- st.session_state.session_id = int(match.group(1))
- try:
- with open(file_path, "r", encoding="utf-8") as f:
- st.session_state.messages = json.load(f)
- except Exception:
- st.session_state.messages = []
- def load_latest_history():
- history_files = sorted(glob.glob(os.path.join(data_dir, "*.json")), key=os.path.getmtime)
- if history_files:
- load_history(history_files[-1])
- if "messages" not in st.session_state:
- st.session_state.messages = []
- load_latest_history()
- if "session_id" not in st.session_state:
- st.session_state.session_id = load_session_id()
- # 用于控制“单块展开”的状态集合(索引集合)
- if "expanded_messages" not in st.session_state:
- st.session_state.expanded_messages = set()
- def save_history():
- file_path = os.path.join(data_dir, f"chat_history_{st.session_state.session_id}.json")
- try:
- json.dumps(st.session_state.messages)
- with open(file_path, "w", encoding="utf-8") as f:
- json.dump(st.session_state.messages, f, ensure_ascii=False)
- except (TypeError, ValueError):
- st.error("数据格式不正确,无法保存记录")
- def move_history(file_path: str):
- if os.path.exists(file_path):
- destination_file = os.path.join(backup_dir, os.path.basename(file_path))
- shutil.move(file_path, destination_file)
- def delete_history(file_path: str):
- if os.path.exists(file_path):
- os.remove(file_path)
- def text_from_content(content):
- if isinstance(content, str):
- return content
- elif isinstance(content, list):
- return " ".join([part.get("text", "") for part in content if part.get("type") == "text"])
- return str(content)
- def get_chat_title(messages):
- if messages:
- title = text_from_content(messages[0].get("content", "")) or "空的聊天"
- return title[:12]
- return "空的聊天"[:12]
- def show_history_files(page: int = 0, page_size: int = 10):
- with st.sidebar.expander("历史聊天记录"):
- history_files = [(f, os.path.getmtime(f)) for f in glob.glob(os.path.join(data_dir, "*.json"))]
- history_files.sort(key=lambda x: x[1], reverse=True)
- total_files = len(history_files)
- total_pages = (total_files // page_size) + (1 if total_files % page_size > 0 else 0)
- # 页码范围校正
- if page >= total_pages and total_pages > 0:
- st.session_state.current_page = max(0, total_pages - 1)
- st.rerun()
- elif page < 0:
- st.session_state.current_page = 0
- st.rerun()
- start = page * page_size
- end = start + page_size
- display_files = history_files[start:end]
- if not display_files:
- st.write("无记录。")
- else:
- for file_path, _mtime in display_files:
- try:
- with open(file_path, "r", encoding="utf-8") as f:
- chat_history = json.load(f)
- except Exception:
- continue
- file_name = get_chat_title(chat_history)
- col1, col2, col3 = st.sidebar.columns([3, 1, 1])
- with col1:
- if st.button(file_name, key=f"load_{os.path.basename(file_path)}"):
- load_history(file_path)
- # 切换历史时,清空单块展开状态
- st.session_state.expanded_messages = set()
- st.rerun()
- with col2:
- if st.button("📦", key=f"move_{os.path.basename(file_path)}", help="移动到备份文件夹"):
- move_history(file_path)
- st.success(f"{file_name} 已移动到备份。")
- st.rerun()
- with col3:
- if st.button("❌", key=f"delete_{os.path.basename(file_path)}", help="删除"):
- delete_history(file_path)
- st.success(f"{file_name} 已删除。")
- st.rerun()
- col_prev, col_next = st.sidebar.columns([1, 1])
- with col_prev:
- if st.button("上一页", disabled=page <= 0, key="history_prev_page"):
- st.session_state.current_page -= 1
- st.rerun()
- with col_next:
- if st.button("下一页", disabled=page >= total_pages - 1, key="history_next_page"):
- st.session_state.current_page += 1
- st.rerun()
- if "current_page" not in st.session_state:
- st.session_state.current_page = 0
- def export_message(content):
- # 文本提取与转义
- if isinstance(content, str):
- processed_content = content.replace('\n', '\n').replace('"', '"')
- elif isinstance(content, list):
- processed_content = " ".join([part.get("text", "") for part in content if part.get("type") == "text"])
- processed_content = processed_content.replace('\n', '\n').replace('"', '"')
- else:
- processed_content = str(content)
- timestamp = datetime.datetime.now().strftime("%m%d%H%M")
- first_10 = (
- processed_content[:10]
- .replace(' ', '')
- .replace('/', '')
- .replace('\\', '')
- .replace(':', '')
- .replace('```', '')
- )
- filename = f"{timestamp}_{first_10}.txt"
- file_path = os.path.join(blog_dir, filename)
- try:
- with open(file_path, "w", encoding="utf-8") as f:
- f.write(processed_content)
- st.success(f"已导出到 {file_path}")
- except Exception as e:
- st.error(f"导出失败:{e}")
- # 侧边栏:输出模式
- output_mode = st.sidebar.selectbox("选择输出模式", ["流式输出 (Stream)", "非流式输出 (Non-stream)"])
- # ========== 高亮样式与收缩样式 ==========
- st.markdown("""
- <style>
- mark.hl {
- background: #fff3a3;
- color: inherit;
- padding: 0 2px;
- border-radius: 3px;
- }
- /* 收缩为 3 行预览 */
- .clamp-5 {
- display: -webkit-box;
- -webkit-line-clamp: 3;
- -webkit-box-orient: vertical;
- overflow: hidden;
- }
- .message-box {
- margin-bottom: 0.25rem;
- }
- /* 让按钮文字不换行(包含导出按钮) */
- .stButton > button {
- white-space: nowrap;
- }
- </style>
- """, unsafe_allow_html=True)
- # 侧边栏:当前对话搜索(仅高亮,不提供“上/下”和“清除定位”按钮)
- def message_contains_query(content, query: str) -> bool:
- if not query:
- return False
- q = query.lower()
- if isinstance(content, str):
- return q in content.lower()
- elif isinstance(content, list):
- for part in content:
- if part.get("type") == "text" and q in part.get("text", "").lower():
- return True
- return False
- else:
- return q in str(content).lower()
- search_query = st.sidebar.text_input("搜索当前对话内容")
- # 统计匹配数量(仅展示数量,不提供跳转)
- if search_query:
- results = [i for i, m in enumerate(st.session_state.messages) if message_contains_query(m.get("content"), search_query)]
- st.sidebar.write(f"共找到 {len(results)} 条匹配。")
- else:
- st.sidebar.write("无匹配。")
- # 历史消息数量控制
- if "history_count" not in st.session_state:
- st.session_state.history_count = 0
- max_history_count = len(st.session_state.messages) if st.session_state.messages else 10
- st.sidebar.slider(
- "选择使用的历史消息数量(共" + str(len(st.session_state.messages)) + "条)",
- min_value=0,
- max_value=max_history_count,
- value=st.session_state.history_count,
- key="history_count"
- )
- st.sidebar.write(f"您选择的历史消息数量是: {st.session_state.history_count}")
- # 文件上传保存为 HTTP 路径
- BASE_URL = "http://117.50.195.224:8007/download/"
- def save_uploaded_file(uploaded_file):
- unique_filename = f"{uuid.uuid4()}_{uploaded_file.name}"
- file_path = os.path.join(upload_dir, unique_filename)
- with open(file_path, "wb") as f:
- f.write(uploaded_file.getvalue())
- return f"{BASE_URL}{unique_filename}"
- # ========== 高亮辅助函数 ==========
- def highlight_text_safe(text: str, query: str) -> str:
- if not query:
- return html.escape(text)
- pattern = re.compile(re.escape(query), re.IGNORECASE)
- parts = pattern.split(text)
- matches = pattern.findall(text)
- out = []
- for i, part in enumerate(parts):
- out.append(html.escape(part))
- if i < len(matches):
- out.append(f"<mark class='hl'>{html.escape(matches[i])}</mark>")
- return "".join(out)
- def render_content(content, query: str = None, clamp_lines: int = None):
- # clamp_lines 为 None 时不收缩;为整数时收缩为该行数(当前实现为 5)
- clamp_class = "clamp-5" if clamp_lines and clamp_lines > 0 else None
- def render_text(txt: str):
- if query:
- html_text = highlight_text_safe(txt, query)
- if clamp_class:
- st.markdown(f"<div class='message-box {clamp_class}'>{html_text}</div>", unsafe_allow_html=True)
- else:
- st.markdown(html_text, unsafe_allow_html=True)
- else:
- if clamp_class:
- st.markdown(f"<div class='message-box {clamp_class}'>{html.escape(txt)}</div>", unsafe_allow_html=True)
- else:
- st.markdown(txt)
- if isinstance(content, str):
- render_text(content)
- elif isinstance(content, list):
- for part in content:
- if part.get("type") == "text":
- txt = part.get("text", "")
- render_text(txt)
- elif part.get("type") == "image_url":
- url = part["image_url"]["url"]
- # 图片不收缩
- if "base64," in url:
- base64_data = url.split("base64,")[1]
- image_bytes = base64.b64decode(base64_data)
- st.image(image_bytes, caption="上传的图片")
- else:
- st.image(url, caption="上传的图片")
- else:
- # 其他类型直接显示
- if clamp_class:
- st.markdown(f"<div class='message-box {clamp_class}'>{html.escape(str(content))}</div>", unsafe_allow_html=True)
- else:
- st.markdown(str(content))
- else:
- if clamp_class:
- st.markdown(f"<div class='message-box {clamp_class}'>{html.escape(str(content))}</div>", unsafe_allow_html=True)
- else:
- st.markdown(str(content))
- # ========== 渲染完整聊天消息(支持高亮与“历史自动收缩”/“搜索时展开”) ==========
- messages_len = len(st.session_state.messages)
- for idx, message in enumerate(st.session_state.messages):
- # 搜索时全部展开;非搜索时,除了最后一条,其余收缩为 5 行
- is_searching = bool(search_query)
- default_clamp = None if is_searching else (5 if idx < messages_len - 1 else None)
- is_expanded = (idx in st.session_state.expanded_messages)
- clamp_lines = None if (is_searching or is_expanded) else default_clamp
- with st.chat_message(message["role"]):
- add_hit = bool(search_query and message_contains_query(message.get("content"), search_query))
- render_content(
- message.get("content"),
- query=search_query if add_hit else None,
- clamp_lines=clamp_lines
- )
- # 同一行显示“展开/收起”和“导出”按钮,并保证导出按钮文本不换行
- spacer_col, export_col, expand_col = st.columns([8, 1, 1])
- with expand_col:
- if default_clamp and not is_searching:
- if not is_expanded:
- if st.button("\>\>", key=f"expand_{idx}"):
- st.session_state.expanded_messages.add(idx)
- st.rerun()
- else:
- if st.button("<<", key=f"collapse_{idx}"):
- if idx in st.session_state.expanded_messages:
- st.session_state.expanded_messages.remove(idx)
- st.rerun()
- with export_col:
- if message["role"] == "assistant":
- export_key = f"export_{idx}"
- if st.button("导出", key=export_key):
- export_message(message.get("content"))
- # 注:根据需求“去掉上,下和清除定位按钮,只高亮即可”,不再自动滚动到定位消息,也不提供定位跳转控件。
- # ========== 处理 chat_input ==========
- prompt_input = st.chat_input(
- "What is up?",
- accept_file="multiple",
- file_type=["jpg", "jpeg", "png", "txt", "pdf", "doc", "docx"]
- )
- if prompt_input:
- text_content = prompt_input.text or ""
- uploaded_files = prompt_input.files
- content = text_content
- additional_prompt = ""
- if uploaded_files:
- content_parts = [{"type": "text", "text": text_content}]
- for file in uploaded_files:
- file_type = (file.type.split('/')[1].lower() if '/' in file.type else file.type.lower()).strip()
- if file_type in ["jpg", "jpeg", "png"]:
- image_bytes = file.getvalue()
- base64_image = base64.b64encode(image_bytes).decode('utf-8')
- content_parts.append({
- "type": "image_url",
- "image_url": {"url": f"data:image/{file_type};base64,{base64_image}"}
- })
- else:
- http_url = save_uploaded_file(file)
- additional_prompt += f"本次提问包含:{http_url} 文件\n"
- content = content_parts if len(content_parts) > 1 else text_content
- if additional_prompt:
- if isinstance(content, list):
- content[0]["text"] = (content[0]["text"] + "\n" + additional_prompt).strip()
- else:
- content = (content + "\n" + additional_prompt).strip()
- # 用户发送新消息后,“历史的自动收缩”:清空展开状态集合
- st.session_state.expanded_messages = set()
- st.session_state.messages.append({"role": "user", "content": content})
- with st.chat_message("user"):
- # 最新消息不收缩
- render_content(content, clamp_lines=None)
- with st.chat_message("assistant"):
- client.api_key = model_keys[selected_model]
- stream = output_mode == "流式输出 (Stream)"
- try:
- if st.session_state.history_count > 0:
- messages_to_send = [
- {"role": m["role"], "content": m["content"]}
- for m in st.session_state.messages[-st.session_state.history_count:]
- ]
- else:
- messages_to_send = [{"role": "user", "content": content}]
- res = client.chat.completions.create(
- model=selected_model,
- messages=messages_to_send,
- stream=stream,
- )
- if stream:
- assistant_message = st.write_stream(res)
- if assistant_message:
- st.session_state.messages.append({"role": "assistant", "content": assistant_message})
- else:
- st.warning("收到空响应。")
- else:
- if len(res.choices) > 0:
- assistant_message = res.choices[0].message.content
- if assistant_message:
- render_content(assistant_message)
- st.session_state.messages.append({"role": "assistant", "content": assistant_message})
- else:
- st.warning("收到空响应。")
- else:
- st.warning("响应格式不正确,未找到有效消息。")
- save_history()
- # 确保新增的最后一条消息在有搜索时也能被高亮与计数,强制重渲染
- st.rerun()
- except Exception as e:
- st.error(f"发生错误:{e}")
- # ========== 侧边栏操作 ==========
- st.sidebar.header("操作")
- if st.sidebar.button("New Chat"):
- st.session_state.messages = []
- st.session_state.session_id = load_session_id()
- st.session_state.session_id += 1
- save_session_id(st.session_state.session_id)
- # 新建会话时重置展开状态
- st.session_state.expanded_messages = set()
- st.success("当前会话已清空。")
- st.sidebar.header("历史对话")
- show_history_files(st.session_state.current_page)
|