将处理IMDB数据的Python函数重构为类的技术咨询
IMDB SQLite 工具类重构指南
重构后完整代码
import argparse import sqlite3 from sqlite3 import Error import pandas as pd class IMDBDatabase: def __init__(self, db_file, table_name): self.db_file = db_file self.table_name = table_name self.conn = self.create_connection() def create_connection(self): """创建SQLite数据库连接""" conn = None try: conn = sqlite3.connect(self.db_file) return conn except Error as e: print(f"数据库连接失败: {e}") return conn def create_table(self, create_table_sql): """根据SQL语句创建表""" if not self.conn: print("未建立数据库连接") return try: c = self.conn.cursor() c.execute(create_table_sql) self.conn.commit() print("表创建成功") except Error as e: print(f"创建表失败: {e}") def seed_data_from_csv(self, csv_path='imdb_top_1000.csv'): """从CSV文件导入数据到表""" if not self.conn: print("未建立数据库连接") return try: df = pd.read_csv(csv_path) df.to_sql(self.table_name, self.conn, if_exists='append', index=False) self.conn.commit() print("数据导入成功") except Exception as e: print(f"数据导入失败: {e}") def get_top10_movies(self): """获取IMDB评分Top10的电影""" if not self.conn: print("未建立数据库连接") return [] try: cur = self.conn.cursor() cur.execute(f""" SELECT Series_Title, Imdb_Rating FROM {self.table_name} ORDER BY Imdb_Rating DESC LIMIT 10 """) rows = cur.fetchall() return rows except Error as e: print(f"查询失败: {e}") return [] def get_top10_lead_actors(self): """获取平均IMDB评分Top10的主演""" if not self.conn: print("未建立数据库连接") return [] try: cur = self.conn.cursor() cur.execute(f""" SELECT Star1, AVG(Imdb_Rating) AS avg_imdb FROM {self.table_name} GROUP BY Star1 ORDER BY avg_imdb DESC LIMIT 10 """) rows = cur.fetchall() return rows except Error as e: print(f"查询失败: {e}") return [] def get_movies_by_year(self, year): """获取指定年份的电影,按IMDB评分排序""" if not self.conn: print("未建立数据库连接") return [] try: cur = self.conn.cursor() cur.execute(f""" SELECT Series_Title, Imdb_Rating FROM {self.table_name} WHERE Released_Year = ? ORDER BY Imdb_Rating DESC """, (year,)) rows = cur.fetchall() return rows except Error as e: print(f"查询失败: {e}") return [] def get_longest_movie_by_year(self, year): """获取指定年份时长最长的电影(转换为小时)""" if not self.conn: print("未建立数据库连接") return [] try: cur = self.conn.cursor() # 提取数字部分转换为分钟,再转为小时 cur.execute(f""" SELECT Series_Title, CAST(REPLACE(Runtime, ' min', '') AS INTEGER)/60.0 AS runtime_hours FROM {self.table_name} WHERE Released_Year = ? ORDER BY CAST(REPLACE(Runtime, ' min', '') AS INTEGER) DESC LIMIT 1 """, (year,)) rows = cur.fetchall() return rows except Error as e: print(f"查询失败: {e}") return [] def get_highest_grossing_year(self): """获取平均票房最高的年份及对应平均票房""" if not self.conn: print("未建立数据库连接") return [] try: cur = self.conn.cursor() cur.execute(f""" SELECT Released_Year, AVG(CAST(REPLACE(Gross, ',', '') AS INTEGER)) AS avg_gross FROM {self.table_name} WHERE Gross IS NOT NULL GROUP BY Released_Year ORDER BY avg_gross DESC LIMIT 1 """) rows = cur.fetchall() return rows except Error as e: print(f"查询失败: {e}") return [] def search_movies_by_name(self, movie_name): """根据电影名称模糊查询相关电影及评分""" if not self.conn: print("未建立数据库连接") return [] try: cur = self.conn.cursor() cur.execute(f""" SELECT Series_Title, Imdb_Rating FROM {self.table_name} WHERE Series_Title LIKE ? ORDER BY Imdb_Rating DESC """, (f'%{movie_name}%',)) rows = cur.fetchall() return rows except Error as e: print(f"查询失败: {e}") return [] def __del__(self): """对象销毁时关闭数据库连接""" if self.conn: self.conn.close() print("数据库连接已关闭") def main(): parser = argparse.ArgumentParser(description="IMDB数据库查询工具") parser.add_argument('--database', type=str, required=True, help="SQLite数据库文件路径") parser.add_argument('--table', type=str, required=True, help="存储电影数据的表名") parser.add_argument('--query', type=str, required=True, choices=[ 'top_10_movies', 'top_10_actors', 'year', 'longest_movie', 'gross_year', 'find_movie' ], help="要执行的查询类型") parser.add_argument('--seed', action='store_true', help="是否从CSV导入数据到数据库") parser.add_argument('--movie', type=str, help="查询电影名称(用于find_movie)") parser.add_argument('--year', type=str, help="查询年份(用于year、longest_movie)") args = parser.parse_args() # 初始化数据库工具类 imdb_db = IMDBDatabase(args.database, args.table) # 导入数据(如果指定--seed) if args.seed: imdb_db.seed_data_from_csv() # 执行指定查询 if args.query == "top_10_movies": print("Top 10高评分电影:") results = imdb_db.get_top10_movies() for row in results: print(f"电影名: {row[0]}, 评分: {row[1]}") elif args.query == "top_10_actors": print("Top 10平均评分主演:") results = imdb_db.get_top10_lead_actors() for row in results: print(f"主演: {row[0]}, 平均评分: {round(row[1], 2)}") elif args.query == "year": if not args.year: print("错误:需要指定--year参数") return print(f"{args.year}年的电影(按评分排序):") results = imdb_db.get_movies_by_year(args.year) if not results: print(f"未找到{args.year}年的电影数据") for row in results: print(f"电影名: {row[0]}, 评分: {row[1]}") elif args.query == "longest_movie": if not args.year: print("错误:需要指定--year参数") return print(f"{args.year}年时长最长的电影:") results = imdb_db.get_longest_movie_by_year(args.year) if results: print(f"电影名: {results[0][0]}, 时长: {round(results[0][1], 2)}小时") else: print(f"未找到{args.year}年的电影数据") elif args.query == "gross_year": print("平均票房最高的年份:") results = imdb_db.get_highest_grossing_year() if results: print(f"年份: {results[0][0]}, 平均票房: {round(results[0][1], 2)}") else: print("未找到票房数据") elif args.query == "find_movie": if not args.movie: print("错误:需要指定--movie参数") return print(f"包含'{args.movie}'的电影:") results = imdb_db.search_movies_by_name(args.movie) if not results: print(f"未找到包含'{args.movie}'的电影") for row in results: print(f"电影名: {row[0]}, 评分: {row[1]}") if __name__ == '__main__': main()
关键重构点说明
- 封装核心资源:将数据库连接、表名由类实例统一管理,避免在每个函数中重复传递参数,代码更简洁易维护
- 修复硬编码问题:原代码中固定的年份、电影名改为方法参数,支持用户通过命令行传入动态值,提升工具灵活性
- 参数化查询优化:除表名外(SQLite不支持表名参数化),其余查询条件均使用参数化写法,避免SQL注入风险
- 资源自动管理:通过
__del__方法自动关闭数据库连接,防止资源泄漏 - 职责分离:查询方法仅负责获取数据,打印逻辑独立放在main函数中,提高方法复用性(可直接返回结果用于其他业务场景)
- 错误处理增强:每个方法添加连接检查和异常捕获,输出明确的错误提示,便于排查问题
- 命令行参数优化:添加
choices限制查询类型,避免无效输入;修复原代码中参数校验错误(如原year查询错误检查movie参数)
内容的提问来源于stack exchange,提问作者Jolene Lester
相关产品推荐
相关产品推荐

