(self, input_dict, mode)
| 448 | return input_dict |
| 449 | |
| 450 | def step_activate_learning(self, input_dict, mode): |
| 451 | query = input_dict["query"] |
| 452 | if self.question_type == QuestionType.EASY.value: |
| 453 | table_info = input_dict["table_info"] |
| 454 | else: |
| 455 | table_info = input_dict["add_table_info"] |
| 456 | init_sql = input_dict["sql_results"]["sql"] |
| 457 | try: |
| 458 | correct_sql_by_case_result = correct_sql_by_case(query, table_info, init_sql) |
| 459 | correct_sql = extract_sql(correct_sql_by_case_result, init_sql) |
| 460 | except Exception as e: |
| 461 | init_sql = "error" |
| 462 | correct_sql_by_case_result = correct_sql_by_case(query, table_info, init_sql) |
| 463 | correct_sql = extract_sql(correct_sql_by_case_result, init_sql) |
| 464 | print(e.__str__()) |
| 465 | if "error" not in correct_sql and "CAST" not in correct_sql: |
| 466 | sql = f"{correct_sql}" |
| 467 | else: |
| 468 | print(correct_sql) |
| 469 | sql = input_dict["init_sql_results"].get("sql", "error") |
| 470 | input_dict["act_sql_results"] = {"sql": sql, "correct_sql_by_case_result": correct_sql_by_case_result} |
| 471 | if mode == "debug": |
| 472 | print("-------error case--------") |
| 473 | print(input_dict["act_sql_results"]) |
| 474 | return input_dict |
| 475 | |
| 476 | def save_sql_txt(self, save_file_name, suffix_name, sql, dataset, db_id): |
| 477 | filename, ext = os.path.splitext(save_file_name) |
no test coverage detected