You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

将处理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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.09 00:40:28