Files

182 lines
6.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""
设置管理器:SQLite settings 表 + config.py 默认值合并
- get_setting(key, default) / set_setting(key, value) / all_settings() / save_all(dict)
- 提供 LLM / 邮件 / 监控 三段配置的便捷读取
"""
import time
from database import query, execute, query_one
from config import LLM_BASE_URL, LLM_API_KEY, LLM_MODEL, LLM_MAX_TOKENS, \
LLM_TEMPERATURE, LLM_TIMEOUT, MAIL_DEFAULTS, MONITOR_DEFAULTS, TRACKING_DEFAULTS
def _truthy(v):
"""容错布尔解析:兼容 1/0/True/False/true/false/on/yes"""
if isinstance(v, bool):
return v
s = str(v or "").strip().lower()
return s in ("1", "true", "yes", "on", "y")
def _bool_str(v):
return "1" if _truthy(v) else "0"
def get_setting(key, default=""):
r = query_one("SELECT value FROM settings WHERE key=?", (key,))
return r["value"] if r else default
def set_setting(key, value):
execute("INSERT OR REPLACE INTO settings(key, value) VALUES(?,?)", (key, str(value)))
def all_settings():
rows = query("SELECT key, value FROM settings")
return {r["key"]: r["value"] for r in rows}
def save_all(pairs):
"""pairs: {key: value},仅保存存在的 key(白名单)"""
for k, v in pairs.items():
set_setting(k, v)
# ===================================================================== LLM
def llm_config():
"""运行时 LLM 配置(设置 > 默认)"""
return {
"base_url": get_setting("llm_base_url", LLM_BASE_URL).rstrip("/"),
"api_key": get_setting("llm_api_key", LLM_API_KEY),
"model": get_setting("llm_model", LLM_MODEL),
"max_tokens": int(get_setting("llm_max_tokens", LLM_MAX_TOKENS)),
"temperature": float(get_setting("llm_temperature", LLM_TEMPERATURE)),
"timeout": int(get_setting("llm_timeout", LLM_TIMEOUT)),
}
# ===================================================================== 邮件
def mail_config():
return {
"smtp_host": get_setting("smtp_host", MAIL_DEFAULTS["smtp_host"]),
"smtp_port": int(get_setting("smtp_port", MAIL_DEFAULTS["smtp_port"])),
"smtp_user": get_setting("smtp_user", MAIL_DEFAULTS["smtp_user"]),
"smtp_pass": get_setting("smtp_pass", MAIL_DEFAULTS["smtp_pass"]),
"smtp_mode": get_setting("smtp_mode", MAIL_DEFAULTS["smtp_mode"]),
"email_to": get_setting("email_to", MAIL_DEFAULTS["email_to"]),
"sender_name": get_setting("sender_name", MAIL_DEFAULTS["sender_name"]),
"email_enabled": _truthy(get_setting("email_enabled", "1")),
}
# ===================================================================== 监控
def monitor_config():
cats = get_setting("monitor_categories", "").strip()
# 关键词去重保序
kws = get_setting("monitor_keywords", MONITOR_DEFAULTS["monitor_keywords"])
kw_list = []
for k in kws.split(","):
k = k.strip()
if k and k not in kw_list:
kw_list.append(k)
return {
"enabled": _truthy(get_setting("monitor_enabled", "1")),
"interval_min": max(5, int(get_setting("monitor_interval", MONITOR_DEFAULTS["monitor_interval"]))),
"categories": [c for c in cats.split(",") if c] if cats else [],
"sentiment_weight": float(get_setting("monitor_sentiment", MONITOR_DEFAULTS["monitor_sentiment"])),
"importance_threshold": float(get_setting("monitor_importance", MONITOR_DEFAULTS["monitor_importance"])),
"keywords": kw_list,
}
def monitor_state():
"""监控运行状态(最近处理 id / 上次扫描时间)"""
return {
"last_news_id": int(get_setting("monitor_last_news_id", "0")),
"last_scan": get_setting("monitor_last_scan", ""),
"last_sent_count": int(get_setting("monitor_last_sent", "0")),
}
def set_monitor_state(last_news_id=None, last_scan=None, last_sent=None):
if last_news_id is not None:
set_setting("monitor_last_news_id", last_news_id)
if last_scan is not None:
set_setting("monitor_last_scan", last_scan)
if last_sent is not None:
set_setting("monitor_last_sent", last_sent)
# ===================================================================== 持仓跟踪
def tracking_config():
return {
"enabled": _truthy(get_setting("tracking_enabled", TRACKING_DEFAULTS["tracking_enabled"])),
"interval_min": max(15, int(get_setting("tracking_interval", TRACKING_DEFAULTS["tracking_interval"]))),
"notify": _truthy(get_setting("tracking_notify", TRACKING_DEFAULTS["tracking_notify"])),
"impact_threshold": float(get_setting("tracking_impact_threshold", TRACKING_DEFAULTS["tracking_impact_threshold"])),
}
# ===================================================================== 静默期
def quiet_config(prefix):
"""某个自动化任务的静默期配置。prefix: monitor / tracking / report"""
return {
"enabled": _truthy(get_setting(f"{prefix}_quiet_enabled", "0")),
"ranges": (get_setting(f"{prefix}_quiet_ranges", "") or "").strip(),
}
def in_quiet_period(ranges_str=None, now=None):
"""判断当前时刻是否落在静默时段内。
ranges 格式:"22:00-07:30"(支持跨午夜)或 "23:00-07:30,12:00-13:00"(多个时段逗号分隔)。
任一时段命中即返回 True;解析失败/空配置返回 False。
"""
ranges_str = (ranges_str or "").strip()
if not ranges_str:
return False
if now is None:
now = time.localtime()
if hasattr(now, "tm_hour"):
cur = now.tm_hour * 60 + now.tm_min
else: # datetime 对象
cur = now.hour * 60 + now.minute
for part in ranges_str.replace("", ",").split(","):
part = part.strip()
if "-" not in part:
continue
try:
a, b = part.split("-", 1)
sh, sm = a.strip().split(":")
eh, em = b.strip().split(":")
start = int(sh) * 60 + int(sm)
end = int(eh) * 60 + int(em)
except Exception:
continue
if start == end:
continue
if start < end:
if start <= cur < end:
return True
else: # 跨天时段,如 23:00-07:30
if cur >= start or cur < end:
return True
return False
def tracking_state():
return {
"last_run": get_setting("tracking_last_run", ""),
"last_stock": get_setting("tracking_last_stock", ""),
"last_alert": int(get_setting("tracking_last_alert", "0")),
}
def set_tracking_state(last_run=None, last_stock=None, last_alert=None):
if last_run is not None:
set_setting("tracking_last_run", last_run)
if last_stock is not None:
set_setting("tracking_last_stock", last_stock)
if last_alert is not None:
set_setting("tracking_last_alert", last_alert)