MCPcopy Create free account
hub / github.com/FlyingFeather/DEA-SQL / step_activate_learning

Method step_activate_learning

gen_sql.py:450–474  ·  view source on GitHub ↗
(self, input_dict, mode)

Source from the content-addressed store, hash-verified

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)

Callers 1

mainMethod · 0.95

Calls 2

correct_sql_by_caseFunction · 0.90
extract_sqlFunction · 0.90

Tested by

no test coverage detected