modify scripts

This commit is contained in:
oscarz
2025-03-24 10:48:35 +08:00
parent 7ded7c5a19
commit 1521ff1fc0
11 changed files with 248 additions and 565 deletions
-65
View File
@@ -1,65 +0,0 @@
import logging
import os
import inspect
from datetime import datetime
from pathlib import Path
# MySQL 配置
db_config = {
'host': 'testdb',
'user': 'root',
'password': 'mysqlpw',
'database': 'stockdb'
}
log_dir_prefix = '../log'
global_share_data_dir = '/root/sharedata'
global_stock_data_dir = '/root/hostdir/stock_data'
# 获取当前文件所在目录
current_dir = os.path.dirname(os.path.abspath(__file__))
# 获取项目根目录(假设当前文件在 src/strategy 下)
project_root = os.path.abspath(os.path.join(current_dir, '..', '..'))
# 获取log目录
def get_log_directory():
"""
获取项目根目录下的 log 目录路径。如果 log 目录不存在,则自动创建。
"""
# 获取当前文件所在目录
current_dir = Path(__file__).resolve().parent
# 找到项目根目录,假设项目根目录下有一个 log 文件夹
project_root = current_dir
while project_root.name != 'src' and project_root != project_root.parent:
project_root = project_root.parent
project_root = project_root.parent # 回到项目根目录
# 确保 log 目录存在
log_dir = project_root / 'log'
log_dir.mkdir(parents=True, exist_ok=True)
return log_dir
def get_caller_filename():
# 获取调用 setup_logging 的脚本文件名
caller_frame = inspect.stack()[2]
caller_filename = os.path.splitext(os.path.basename(caller_frame.filename))[0]
return caller_filename
# 设置日志配置
def setup_logging(log_filename=None):
# 如果未传入 log_filename,则使用当前脚本名称作为日志文件名
if log_filename is None:
caller_filename = get_caller_filename()
common_log_dir = get_log_directory()
current_date = datetime.now().strftime('%Y%m%d')
# 拼接 log 文件名,将日期加在扩展名前
log_filename = f'{common_log_dir}/{caller_filename}_{current_date}.log'
logging.basicConfig(level=logging.INFO, format='%(asctime)s %(levelname)s [%(filename)s:%(lineno)d] (%(funcName)s) - %(message)s',
handlers=[
logging.FileHandler(log_filename),
logging.StreamHandler()
])
+7 -2
View File
@@ -9,8 +9,13 @@ db_config = {
'database': 'stockdb'
}
global_share_data_dir = '/root/sharedata'
global_stock_data_dir = '/root/hostdir/stock_data'
home_dir = os.path.expanduser("~")
global_host_data_dir = f'{home_dir}/hostdir/stock_data'
global_share_db_dir = f'{home_dir}/sharedata/sqlite'
# 兼容以前的定义
global_stock_data_dir = global_host_data_dir
global_share_data_dir = f'{home_dir}/sharedata'
# 获取当前文件所在目录
current_dir = Path(__file__).resolve().parent
+314
View File
@@ -0,0 +1,314 @@
import os
import json
import requests
import time
import logging
from bs4 import BeautifulSoup
# 获取个股研报列表的指定页
def fetch_reports_by_stock(page_no, start_date="2023-03-10", end_date="2025-03-10", page_size=50, max_retries = 3):
# 请求头
HEADERS = {
"Accept": "application/json, text/javascript, */*; q=0.01",
"Content-Type": "application/json",
"Origin": "https://data.eastmoney.com",
"Referer": "https://data.eastmoney.com/report/stock.jshtml",
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36",
}
# 请求 URL
URL = "https://reportapi.eastmoney.com/report/list2"
payload = {
"beginTime": start_date,
"endTime": end_date,
"industryCode": "*",
"ratingChange": None,
"rating": None,
"orgCode": None,
"code": "*",
"rcode": "",
"pageSize": page_size,
"p": page_no,
"pageNo": page_no,
"pageNum": page_no,
"pageNumber": page_no
}
logging.debug(f'begin: {start_date}, end: {end_date}')
for attempt in range(max_retries):
try:
response = requests.post(URL, headers=HEADERS, json=payload, timeout=10)
response.raise_for_status()
data = response.json()
return data
except requests.RequestException as e:
logging.warning(f"network error on {URL}: {e}, Retring...")
logging.error(f'Fetching failed after max retries. {URL}')
return None # 达到最大重试次数仍然失败
# 获取行业研报列表的指定页
def fetch_reports_by_industry(page_no, start_date="2023-03-10", end_date="2025-03-10", page_size=50, max_retries = 3):
headers = {
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36",
"Referer": "https://data.eastmoney.com/report/industry.jshtml"
}
url = "https://reportapi.eastmoney.com/report/list"
params = {
"cb": "datatable1413600",
"industryCode": "*",
"pageSize": page_size,
"industry": "*",
"rating": "*",
"ratingChange": "*",
"beginTime": start_date,
"endTime": end_date,
"pageNo": page_no,
"fields": "",
"qType": 1,
"orgCode": "",
"rcode": "",
"p": page_no,
"pageNum": page_no,
"pageNumber": page_no,
"_": int(time.time() * 1000) # 动态时间戳
}
for attempt in range(max_retries):
try:
response = requests.get(url, headers=headers, params=params, timeout=10)
response.raise_for_status()
# 去掉回调函数包装
json_text = response.text.strip("datatable1413600(").rstrip(");")
data = json.loads(json_text)
return data
except requests.RequestException as e:
logging.warning(f"network error on {url}: {e}, Retring...")
return None
except json.JSONDecodeError as e:
logging.warning(f"json decode error on {url}: {e}, Retring...")
return None
logging.error(f'Fetching failed after max retries. {url}')
return None # 达到最大重试次数仍然失败
# 获取宏观研报列表的指定页
def fetch_reports_by_macresearch(page_no, start_date="2023-03-10", end_date="2025-03-10", page_size=50, max_retries = 3):
headers = {
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36",
"Referer": "https://data.eastmoney.com/report/macresearch.jshtml"
}
url = "https://reportapi.eastmoney.com/report/jg"
params = {
"cb": "datatable2612129",
"industryCode": "*",
"pageSize": page_size,
"author": "",
"beginTime": start_date,
"endTime": end_date,
"pageNo": page_no,
"fields": "",
"qType": 3,
"orgCode": "",
"rcode": "",
"p": page_no,
"pageNum": page_no,
"pageNumber": page_no,
"_": int(time.time() * 1000) # 动态时间戳
}
for attempt in range(max_retries):
try:
response = requests.get(url, headers=headers, params=params, timeout=10)
response.raise_for_status()
# 去掉回调函数包装
json_text = response.text.strip("datatable2612129(").rstrip(");")
data = json.loads(json_text)
return data
except requests.RequestException as e:
logging.warning(f"network error on {url}: {e}, Retring...")
return None
except json.JSONDecodeError as e:
logging.warning(f"json decode error on {url}: {e}, Retring...")
return None
logging.error(f'Fetching failed after max retries. {url}')
return None # 达到最大重试次数仍然失败
# 获取策略研报列表的指定页
def fetch_reports_by_strategy(page_no, start_date="2023-03-10", end_date="2025-03-10", page_size=50, max_retries = 3):
headers = {
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36",
"Referer": "https://data.eastmoney.com/report/strategyreport.jshtml"
}
url = "https://reportapi.eastmoney.com/report/jg"
params = {
"cb": "datatable5349866",
"industryCode": "*",
"pageSize": page_size,
"author": "",
"beginTime": start_date,
"endTime": end_date,
"pageNo": page_no,
"fields": "",
"qType": 2,
"orgCode": "",
"rcode": "",
"p": page_no,
"pageNum": page_no,
"pageNumber": page_no,
"_": int(time.time() * 1000) # 动态时间戳
}
for attempt in range(max_retries):
try:
response = requests.get(url, headers=headers, params=params, timeout=10)
response.raise_for_status()
# 去掉回调函数包装
json_text = response.text.strip("datatable5349866(").rstrip(");")
data = json.loads(json_text)
return data
except requests.RequestException as e:
logging.warning(f"network error on {url}: {e}, Retring...")
return None
except json.JSONDecodeError as e:
logging.warning(f"json decode error on {url}: {e}, Retring...")
return None
logging.error(f'Fetching failed after max retries. {url}')
return None # 达到最大重试次数仍然失败
# 获取新股研报列表的指定页
def fetch_reports_by_newstock(page_no, start_date="2023-03-10", end_date="2025-03-10", page_size=50, max_retries = 3):
headers = {
"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36",
"Referer": "https://data.eastmoney.com/report/newstock.jshtml"
}
url = "https://reportapi.eastmoney.com/report/newStockList"
params = {
"cb": "datatable5144183",
"pageSize": page_size,
"author": "",
"beginTime": start_date,
"endTime": end_date,
"pageNo": page_no,
"fields": "",
"qType": 4,
"orgCode": "",
"rcode": "",
"p": page_no,
"pageNum": page_no,
"pageNumber": page_no,
"_": int(time.time() * 1000) # 动态时间戳
}
for attempt in range(max_retries):
try:
response = requests.get(url, headers=headers, params=params, timeout=10)
response.raise_for_status()
# 去掉回调函数包装
json_text = response.text.strip("datatable5144183(").rstrip(");")
data = json.loads(json_text)
return data
except requests.RequestException as e:
logging.warning(f"network error on {url}: {e}, Retring...")
return None
except json.JSONDecodeError as e:
logging.warning(f"json decode error on {url}: {e}, Retring...")
return None
logging.error(f'Fetching failed after max retries. {url}')
return None # 达到最大重试次数仍然失败
# 访问指定 infoCode 的页面,提取 PDF 下载链接
def fetch_pdf_link(url, max_retries = 3):
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/113.0.0.0 Safari/537.36"
}
for attempt in range(max_retries):
try:
response = requests.get(url, headers=headers, timeout=10)
response.raise_for_status()
# 解析 HTML
soup = BeautifulSoup(response.text, "html.parser")
pdf_link = soup.find("a", class_="pdf-link")
if pdf_link and "href" in pdf_link.attrs:
return pdf_link["href"]
else:
logging.warning(f"未找到 PDF 链接: {url}")
return None
except requests.RequestException as e:
logging.error(f"请求失败: {url} {e}")
logging.error(f'Fetching failed after max retries. {url}')
return None # 达到最大重试次数仍然失败
def is_valid_pdf(file_path):
try:
with open(file_path, "rb") as f:
header = f.read(4)
return header == b"%PDF"
except Exception as e:
logging.error(f"验证 PDF 失败: {e}")
return False
def download_pdf_wget(pdf_url, save_path):
cmd = f'wget -O "{save_path}" "{pdf_url}" --quiet --user-agent="Mozilla/5.0"'
os.system(cmd)
return os.path.exists(save_path) and is_valid_pdf(save_path)
# 下载 PDF 并保存到本地
def download_pdf(pdf_url, save_path, max_retries=5):
for attempt in range(max_retries):
down = download_pdf_wget(pdf_url, save_path)
if down:
return True
return False
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/113.0.0.0 Safari/537.36"
}
try:
response = requests.get(pdf_url, headers=headers, stream=True, timeout=20)
response.raise_for_status()
with open(save_path, "wb") as file:
for chunk in response.iter_content(chunk_size=1024):
file.write(chunk)
return True
except requests.RequestException as e:
logging.error(f"PDF 下载失败: {e}")
return False
+123
View File
@@ -0,0 +1,123 @@
import sqlite3
import logging
import sys
from datetime import datetime
class DatabaseConnectionError(Exception):
pass
class StockReportDB:
# 定义类属性(静态变量)
TBL_STOCK = 'reports_stock'
TBL_NEW_STOCK = 'reports_newstrock'
TBL_STRATEGY = 'reports_strategy'
TBL_MACRESEARCH = 'reports_macresearch'
TBL_INDUSTRY = 'reports_industry'
def __init__(self, db_path):
self.DB_PATH = db_path
self.conn = None
self.cursor = None
try:
self.conn = sqlite3.connect(self.DB_PATH)
self.cursor = self.conn.cursor()
except sqlite3.Error as e:
logging.error(f"数据库连接失败: {e}")
raise DatabaseConnectionError("数据库连接失败")
def __get_table_columns_and_defaults(self, tbl_name):
try:
self.cursor.execute(f"PRAGMA table_info({tbl_name})")
columns = self.cursor.fetchall()
column_info = {}
for col in columns:
col_name = col[1]
default_value = col[4]
column_info[col_name] = default_value
return column_info
except sqlite3.Error as e:
logging.error(f"Error getting table columns: {e}")
return None
def __check_and_process_data(self, data, tbl_name):
column_info = self.__get_table_columns_and_defaults(tbl_name=tbl_name)
if column_info is None:
return None
processed_data = {}
for col, default in column_info.items():
if col == 'id': # 自增主键,不需要用户提供
continue
if col == 'created_at' or col == 'updated_at': # 日期函数,用户自己指定即可
continue
if col in ['author', 'authorID']:
values = data.get(col, [])
processed_data[col] = ','.join(values)
elif col in data:
processed_data[col] = data[col]
else:
if default is not None:
processed_data[col] = default
else:
processed_data[col] = None
return processed_data
def insert_or_update_common(self, data, tbl_name, uniq_key='infoCode'):
try:
processed_data = self.__check_and_process_data(data, tbl_name)
if processed_data is None:
return None
columns = ', '.join(processed_data.keys())
values = list(processed_data.values())
placeholders = ', '.join(['?' for _ in values])
update_clause = ', '.join([f"{col}=EXCLUDED.{col}" for col in processed_data.keys() if col != 'infoCode']) + ', updated_at=datetime(\'now\', \'localtime\')'
sql = f'''
INSERT INTO {tbl_name} ({columns}, updated_at)
VALUES ({placeholders}, datetime('now', 'localtime'))
ON CONFLICT (infoCode) DO UPDATE SET {update_clause}
'''
self.cursor.execute(sql, values)
self.conn.commit()
# 获取插入或更新后的 report_id
self.cursor.execute(f"SELECT id FROM {tbl_name} WHERE {uniq_key} = ?", (data["infoCode"],))
report_id = self.cursor.fetchone()[0]
return report_id
except sqlite3.Error as e:
logging.error(f"Error inserting or updating data: {e}")
return None
def query_reports_comm(self, tbl_name, querystr='', limit=None):
try:
if tbl_name in [StockReportDB.TBL_STOCK, StockReportDB.TBL_NEW_STOCK, StockReportDB.TBL_INDUSTRY, StockReportDB.TBL_MACRESEARCH, StockReportDB.TBL_STRATEGY]:
sql = f"SELECT id, infoCode, title, orgSName, industryName, stockName, publishDate FROM {tbl_name} WHERE 1=1 {querystr}"
else:
logging.warning(f'wrong table name: {tbl_name}')
return None
if limit:
sql = sql + f' limit {limit}'
self.cursor.execute(sql)
results = self.cursor.fetchall()
# 获取列名
column_names = [description[0] for description in self.cursor.description]
# 将结果转换为字典列表
result_dict_list = []
for row in results:
row_dict = {column_names[i]: value for i, value in enumerate(row)}
result_dict_list.append(row_dict)
return result_dict_list
except sqlite3.Error as e:
logging.error(f"查询 href 失败: {e}")
return None
def __del__(self):
if self.conn:
self.conn.close()
+333
View File
@@ -0,0 +1,333 @@
import json
import time
import csv
import os
import re
import argparse
import shutil
import logging
from datetime import datetime, timedelta
from functools import partial
import src.crawler.em.reports as em
import src.utils.utils as utils
from src.config.config import global_host_data_dir, global_share_db_dir
from src.db_utils.reports import StockReportDB, DatabaseConnectionError
from src.logger.logger import setup_logging
# 初始化日志
setup_logging()
debug = False
force = False
pdf_base_dir = f"{global_host_data_dir}/pdfs" # 下载 PDF 存放目录
# 定义下载页面的链接
map_pdf_page = {
StockReportDB.TBL_STOCK : "https://data.eastmoney.com/report/info/{}.html",
StockReportDB.TBL_NEW_STOCK : "https://data.eastmoney.com/report/info/{}.html",
StockReportDB.TBL_STRATEGY : "https://data.eastmoney.com/report/zw_strategy.jshtml?encodeUrl={}",
StockReportDB.TBL_MACRESEARCH : "https://data.eastmoney.com/report/zw_macresearch.jshtml?encodeUrl={}",
StockReportDB.TBL_INDUSTRY : "https://data.eastmoney.com/report/zw_industry.jshtml?infocode={}"
}
# 定义表名的映射,作为存储路径用
map_tbl_name = {
StockReportDB.TBL_STOCK : '个股研报',
StockReportDB.TBL_NEW_STOCK : '新股研报',
StockReportDB.TBL_STRATEGY : '策略报告',
StockReportDB.TBL_MACRESEARCH : '宏观研究',
StockReportDB.TBL_INDUSTRY : '行业研报'
}
# 初始化数据库连接
db_path = f"{global_share_db_dir}/stock_report.db"
db_tools = None
current_date = datetime.now()
seven_days_ago = current_date - timedelta(days=7)
two_years_ago = current_date - timedelta(days=2*365)
start_date = two_years_ago.strftime("%Y-%m-%d")
end_date = current_date.strftime("%Y-%m-%d")
this_week_date = seven_days_ago.strftime("%Y-%m-%d")
def fetch_reports_list_general(fetch_func, table_name, s_date, e_date, data_dir_prefix):
# 示例:获取前 3 页的数据
max_pages = 100000
page = 1
while page <= max_pages:
while True:
data = fetch_func(page_no=page, start_date=s_date, end_date=e_date, page_size=100)
if data:
break
if page == 1:
max_pages = data.get('TotalPage', 1000000)
for row in data.get('data', []):
# 统一以 infoCode 为 UNIQE 键,所以这里对它进行赋值
if row.get('infoCode') is None and row.get('encodeUrl'):
row['infoCode'] = row['encodeUrl']
row_id = db_tools.insert_or_update_common(row, table_name)
if row_id:
logging.debug(f'insert one row. rowid:{row_id}, ')
else:
logging.warning(f'insert data failed. page : {page}')
return False
# 写本地json文件,必要性不大
#utils.save_json_to_file(data, f'{utils.json_data_dir}/{data_dir_prefix}', f'{data_dir_prefix}_report_{page}.json')
logging.info(f"{page} 页, 获取 {len(data['data'])} 条数据, 共 {max_pages}")
page += 1
time.sleep(1) # 避免请求过快
# 股票所用的url
def parse_func_general(row, tbl_name):
info_code = row['infoCode']
title = row['title'].replace("/", "_").replace("\\", "_")
org_sname = row['orgSName']
stock_name = row['stockName']
industry_name = row['industryName']
publish_date = row['publishDate'].split(" ")[0]
# 建表的时候默认值有点问题
if stock_name == '' or stock_name=="''":
stock_name = 'None'
if industry_name == '':
industry_name = 'None'
if org_sname == '':
org_sname = 'None'
report_type = map_tbl_name.get(tbl_name, 'None')
file_name = f"{publish_date}_{report_type}_{org_sname}_{industry_name}_{stock_name}_{title}.pdf"
url = map_pdf_page.get(tbl_name, None)
if url is None:
logging.warning(f'wrong table name: {tbl_name}')
return None, None
url = url.format(info_code)
# 拼目录
dir_year = publish_date[:4] if len(publish_date)>4 else ''
dir_path = f'{pdf_base_dir}/{dir_year}/{map_tbl_name[tbl_name]}'
os.makedirs(dir_path, exist_ok=True)
return url, os.path.join(dir_path, file_name)
# 通用下载函数
def download_pdf_stock_general(parse_func, tbl_name, querystr='', s_date=start_date, e_date=end_date, limit=None):
# 下载pdf
if s_date:
querystr += f" AND publishDate >= '{s_date} 00:00:00.000' "
if e_date:
querystr += f" AND publishDate <= '{e_date} 23:59:59.999' "
rows = db_tools.query_reports_comm(tbl_name, querystr=querystr, limit=limit)
if rows is None:
rows = []
for row in rows:
url, file_path = parse_func(row, tbl_name)
if url is None or file_path is None:
logging.warning(f'wrong url or file_path. tbl_name: {tbl_name}')
continue
# 已经存在的,跳过
if file_path and os.path.isfile(file_path):
logging.info(f'{file_path} already exists. skipping...')
continue
# 获取pdf链接地址
pdf_url = em.fetch_pdf_link(url)
if pdf_url:
# 下载 PDF
down = em.download_pdf(pdf_url, file_path)
if down:
logging.info(f'saved file {file_path}')
else:
logging.warning(f'download pdf file error. file_path: {pdf_url}, save_path: {file_path}')
else:
logging.warning(f'cannot get pdf link. url: {url}, save_path: {file_path}')
time.sleep(1) # 避免请求过快
# 获取股票报告列表
def fetch_reports_list_stock(s_date=start_date, e_date=end_date):
return fetch_reports_list_general(em.fetch_reports_by_stock, StockReportDB.TBL_STOCK, s_date, e_date, 'stock')
# 获取股票报告列表
def fetch_reports_list_newstock(s_date=start_date, e_date=end_date):
return fetch_reports_list_general(em.fetch_reports_by_newstock, StockReportDB.TBL_NEW_STOCK, s_date, e_date, 'new')
# 获取行业报告列表
def fetch_reports_list_industry(s_date=start_date, e_date=end_date):
return fetch_reports_list_general(em.fetch_reports_by_industry, StockReportDB.TBL_INDUSTRY, s_date, e_date, 'industry')
# 获取行业报告列表
def fetch_reports_list_macresearch(s_date=start_date, e_date=end_date):
return fetch_reports_list_general(em.fetch_reports_by_macresearch, StockReportDB.TBL_MACRESEARCH, s_date, e_date, 'macresearch')
# 获取行业报告列表
def fetch_reports_list_strategy(s_date=start_date, e_date=end_date):
return fetch_reports_list_general(em.fetch_reports_by_strategy, StockReportDB.TBL_STRATEGY, s_date, e_date, 'strategy')
# 下载股票pdf
def download_pdf_stock(s_date=start_date, e_date=end_date):
download_pdf_stock_general(parse_func_general, StockReportDB.TBL_STOCK, ' ', s_date, e_date, limit=2 if debug else None)
def download_pdf_newstock(s_date=start_date, e_date=end_date):
download_pdf_stock_general(parse_func_general, StockReportDB.TBL_NEW_STOCK, ' ', s_date, e_date, limit=2 if debug else None)
def download_pdf_industry(s_date=start_date, e_date=end_date):
download_pdf_stock_general(parse_func_general, StockReportDB.TBL_INDUSTRY, ' ', s_date, e_date, limit=2 if debug else None)
def download_pdf_macresearch(s_date=start_date, e_date=end_date):
download_pdf_stock_general(parse_func_general, StockReportDB.TBL_MACRESEARCH, ' ', s_date, e_date, limit=2 if debug else None)
def download_pdf_strategy(s_date=start_date, e_date=end_date):
download_pdf_stock_general(parse_func_general, StockReportDB.TBL_STRATEGY, ' ', s_date, e_date, limit=2 if debug else None)
# 建立缩写到函数的映射
function_list_map = {
'stock' : fetch_reports_list_stock,
'new' : fetch_reports_list_newstock,
'indust' : fetch_reports_list_industry,
'macro' : fetch_reports_list_macresearch,
'stra' : fetch_reports_list_strategy,
}
function_down_map = {
'stock' : download_pdf_stock,
'new' : download_pdf_newstock,
'indust' : download_pdf_industry,
'macro' : download_pdf_macresearch,
'stra' : download_pdf_strategy,
}
# 获取最新一周的报告列表
def create_last_week_links(s_date=start_date, e_date=end_date):
last_week_dir = os.path.join(pdf_base_dir, 'last_week')
# 如果 last_week 目录存在,先删除它
if os.path.exists(last_week_dir):
for root, dirs, files in os.walk(last_week_dir, topdown=False):
for file in files:
file_path = os.path.join(root, file)
os.remove(file_path)
for dir in dirs:
dir_path = os.path.join(root, dir)
os.rmdir(dir_path)
os.rmdir(last_week_dir)
os.makedirs(last_week_dir)
for root, dirs, files in os.walk(pdf_base_dir):
# 跳过 last_week 目录及其子目录
if 'last_week' in dirs:
dirs.remove('last_week')
for file in files:
if file.endswith('.pdf'):
match = re.match(r'(\d{4}-\d{2}-\d{2})_(.*)\.pdf', file)
if match:
date_str = match.group(1)
if utils.is_within_last_week(date_str):
file_path = os.path.join(root, file)
# 获取子目录名称
sub_dir_name = os.path.basename(os.path.dirname(file_path))
# 生成新的链接名称,添加子目录名前缀
new_file_name = f"[{sub_dir_name}]_{file}"
link_name = os.path.join(last_week_dir, new_file_name)
if not os.path.exists(link_name):
os.symlink(file_path, link_name)
# 执行功能函数
def run_func(function_names, function_map):
global start_date
global end_date
for short_name in function_names:
func = function_map.get(short_name.strip()) # 从映射中获取对应的函数
if callable(func):
#db_tools.update_task_log(task_id, task_status=f'Running {func}')
logging.info(f'exec function: {func}, begin: {start_date}, end: {end_date}')
func(start_date, end_date)
else:
logging.warning(f"Warning: {short_name} is not a valid function shortcut.")
# 主函数
def main(cmd, mode, args_debug, args_force, begin, end):
global debug
debug = args_debug
global force
force = args_force
global start_date
start_date = begin if begin else start_date
global end_date
end_date = end if end else end_date
# 初始化DB
global db_tools
try:
db_tools = StockReportDB(db_path)
# 进行数据库操作
except DatabaseConnectionError as e:
logging.error(f"数据库连接失败: {e}")
return False
# 开启任务
#task_id = db_tools.insert_task_log()
task_id = 0
if task_id is None:
logging.warning(f'insert task log error.')
return None
logging.info(f'running task. id: {task_id}, debug: {debug}, force: {force}, cmd: {cmd}, mode: {mode}')
# 如果是lastweek,我们先执行列表,再执行下载
function_list = []
if mode == 'fetch':
function_list.append(function_list_map)
elif mode == 'down':
function_list.append(function_down_map)
elif mode == 'lastweek':
start_date = this_week_date
function_list.append(function_list_map)
function_list.append(function_down_map)
else:
function_list.append(function_list_map)
# 执行指定的函数
if cmd and mode !='lastweek':
function_names = args.cmd.split(",") # 拆分输入
else:
function_names = function_list_map.keys()
# 遍历功能函数,执行
for function_map in function_list:
run_func(function_names, function_map)
logging.info(f'all process completed!')
#db_tools.finalize_task_log(task_id)
if __name__ == "__main__":
# 命令行参数处理
keys_str = ",".join(function_list_map.keys())
parser = argparse.ArgumentParser(description='fetch iafd data.')
parser.add_argument("--cmd", type=str, help=f"Comma-separated list of function shortcuts: {keys_str}")
parser.add_argument("--mode", type=str, help=f"Fetch list or Download pdf: (fetch, down, lastweek)")
parser.add_argument("--begin", type=str, help=f"begin date, YYYY-mm-dd")
parser.add_argument("--end", type=str, help=f"end date, YYYY-mm-dd")
parser.add_argument('--debug', action='store_true', help='Enable debug mode (limit records)')
parser.add_argument('--force', action='store_true', help='force update (true for rewrite all)')
args = parser.parse_args()
main(args.cmd, args.mode, args.debug, args.force, args.begin, args.end)
+62 -6
View File
@@ -1,10 +1,46 @@
import logging
import os
import inspect
import time
from datetime import datetime
from pathlib import Path
from logging.handlers import RotatingFileHandler
from collections import defaultdict
from src.config.config import get_log_directory, get_src_directory
# 统计日志频率
log_count = defaultdict(int) # 记录日志的次数
last_log_time = defaultdict(float) # 记录上次写入的时间戳
class RateLimitFilter(logging.Filter):
"""
频率限制过滤器:
1. 在 60 秒内,同样的日志最多写入 60 次,超过则忽略
2. 如果日志速率超过 100 条/秒,发出告警
"""
LOG_LIMIT = 600 # 每分钟最多记录相同消息 10 次
def filter(self, record):
global log_count, last_log_time
message_key = record.getMessage() # 获取日志内容
# 计算当前时间
now = time.time()
elapsed = now - last_log_time[message_key]
# 限制相同日志的写入频率
if elapsed < 60: # 60 秒内
log_count[message_key] += 1
if log_count[message_key] > self.LOG_LIMIT:
return False # 直接丢弃
else:
log_count[message_key] = 1 # 超过 60 秒,重新计数
last_log_time[message_key] = now
return True # 允许写入日志
def get_caller_filename():
# 获取调用栈
stack = inspect.stack()
@@ -26,7 +62,7 @@ def get_caller_filename():
return os.path.splitext(os.path.basename(frame_info.filename))[0]
return None
# 设置日志配置
def setup_logging(log_filename=None):
# 如果未传入 log_filename,则使用当前脚本名称作为日志文件名
if log_filename is None:
@@ -35,9 +71,29 @@ def setup_logging(log_filename=None):
current_date = datetime.now().strftime('%Y%m%d')
# 拼接 log 文件名,将日期加在扩展名前
log_filename = f'{common_log_dir}/{caller_filename}_{current_date}.log'
max_log_size = 100 * 1024 * 1024 # 10 MB
max_log_files = 10 # 最多保留 10 个日志文件
file_handler = RotatingFileHandler(log_filename, maxBytes=max_log_size, backupCount=max_log_files)
file_handler.setFormatter(logging.Formatter(
'%(asctime)s %(levelname)s [%(filename)s:%(lineno)d] (%(funcName)s) - %(message)s'
))
console_handler = logging.StreamHandler()
console_handler.setFormatter(logging.Formatter(
'%(asctime)s %(levelname)s [%(filename)s:%(lineno)d] (%(funcName)s) - %(message)s'
))
# 创建 logger
logger = logging.getLogger()
logger.setLevel(logging.INFO)
logger.handlers = [] # 避免重复添加 handler
logger.addHandler(file_handler)
logger.addHandler(console_handler)
# 添加频率限制
rate_limit_filter = RateLimitFilter()
file_handler.addFilter(rate_limit_filter)
console_handler.addFilter(rate_limit_filter)
logging.basicConfig(level=logging.INFO, format='%(asctime)s %(levelname)s [%(filename)s:%(lineno)d] (%(funcName)s) - %(message)s',
handlers=[
logging.FileHandler(log_filename),
logging.StreamHandler()
])
+12 -5
View File
@@ -1,6 +1,7 @@
import pandas as pd
import numpy as np
import os
import warnings
from src.strategy.prepare import fetch_his_kline
import src.config.config as config
import src.crawler.em.stock as em_stock
@@ -15,10 +16,15 @@ def select_stocks(stock_map):
df = fetch_his_kline(stock_code)
close_prices = df['close'].values
# 使用 MyTT 库计算不同周期的 RSI
rsi_6 = RSI(close_prices, 6)
rsi_12 = RSI(close_prices, 12)
rsi_24 = RSI(close_prices, 24)
# 捕获 RuntimeWarning 警告
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always", RuntimeWarning)
# 使用 MyTT 库计算不同周期的 RSI
rsi_6 = RSI(close_prices, 6)
rsi_12 = RSI(close_prices, 12)
rsi_24 = RSI(close_prices, 24)
if w:
print(f"股票代码 {stock_code} {stock_name} 在计算 RSI 时出现警告: {w[0].message}")
df['rsi_6'] = rsi_6
df['rsi_12'] = rsi_12
@@ -96,7 +102,8 @@ if __name__ == "__main__":
codes = ['105.QFIN', '105.FUTU']
# 从网络上获取
stock_map = em_stock.code_by_fs('hk_famous', em_stock.em_market_fs_types['hk_famous'])
plat_id = 'cn_hs300'
stock_map = em_stock.code_by_fs(plat_id, em_stock.em_market_fs_types[plat_id])
if stock_map:
select_stocks(stock_map)
+27
View File
@@ -0,0 +1,27 @@
import re
import os
import json
import time
import csv
import logging
from datetime import datetime
# 保存 JSON 数据到本地文件
def save_json_to_file(data, file_path, file_name):
os.makedirs(file_path, exist_ok=True)
full_name = f"{file_path}/{file_name}"
with open(full_name, "w", encoding="utf-8") as file:
json.dump(data, file, ensure_ascii=False, indent=4)
logging.debug(f"saved json data to: {full_name}")
# 判断日期字符串是否在最近七天内
def is_within_last_week(date_str):
try:
file_date = datetime.strptime(date_str, '%Y-%m-%d')
current_date = datetime.now()
diff = current_date - file_date
return diff.days <= 7
except ValueError:
return False