diff --git a/platform/CTFd/plugins/sql_challenges/README.md b/platform/CTFd/plugins/sql_challenges/README.md index 59a627b3ac..958cd3da2f 100644 --- a/platform/CTFd/plugins/sql_challenges/README.md +++ b/platform/CTFd/plugins/sql_challenges/README.md @@ -64,6 +64,8 @@ scripts/test-sql-judge ### Creating a SQL Challenge +Init scripts exported from a local MySQL can keep their `DROP SCHEMA`, `CREATE SCHEMA` and `USE` statements: the judge maps that schema name onto the execution's temporary database and skips those statements, so `kbo.PLAYER` in init or in a submission resolves to the temporary database. `SET` statements in the init script, such as `SET SQL_MODE='TRADITIONAL'`, stay in effect for the graded statement exactly as they would in one local MySQL session. Leading comment lines, including `-----` separators that MySQL itself rejects, are removed from each statement. Table names are case-insensitive (`lower_case_table_names=1`, the Windows and macOS default), so `Salaries` and `salaries` refer to the same table. This setting is fixed when the MySQL data directory is initialized; an existing `mysql-judge-data` volume must be removed before changing it. + 1. Go to Admin Panel → Challenges → Create Challenge 2. Select "sql" as the challenge type 3. Fill in the challenge details: diff --git a/platform/CTFd/plugins/sql_challenges/docker-compose.integration.yml b/platform/CTFd/plugins/sql_challenges/docker-compose.integration.yml index 3b6d49f940..d7ddaf28f1 100644 --- a/platform/CTFd/plugins/sql_challenges/docker-compose.integration.yml +++ b/platform/CTFd/plugins/sql_challenges/docker-compose.integration.yml @@ -9,6 +9,7 @@ services: - --collation-server=utf8mb4_0900_ai_ci - --disable-log-bin - --event-scheduler=DISABLED + - --lower-case-table-names=1 - --innodb-buffer-pool-size=128M - --innodb-flush-log-at-trx-commit=2 - --max-allowed-packet=16M diff --git a/platform/CTFd/plugins/sql_challenges/sql_judge_integration_test.go b/platform/CTFd/plugins/sql_challenges/sql_judge_integration_test.go index 7a8a40eba2..64a19df390 100644 --- a/platform/CTFd/plugins/sql_challenges/sql_judge_integration_test.go +++ b/platform/CTFd/plugins/sql_challenges/sql_judge_integration_test.go @@ -177,6 +177,50 @@ func TestIntegrationGradedStatementIsReadOnly(t *testing.T) { assertNoTemporaryResources(t, server) } +func TestIntegrationInitSchemaTemplateAndSessionMode(t *testing.T) { + server := integrationServer(t) + init := []string{`-- KBO Database Schema +SET @OLD_UNIQUE_CHECKS=@@UNIQUE_CHECKS, UNIQUE_CHECKS=0; +SET @OLD_SQL_MODE=@@SQL_MODE, SQL_MODE='TRADITIONAL'; +/* schema management from a local export */ +DROP SCHEMA IF EXISTS kbo; +CREATE SCHEMA kbo; +USE kbo; +CREATE TABLE kbo.PLAYER (id INT PRIMARY KEY, team CHAR(3), height INT); +INSERT INTO PLAYER VALUES (1, 'K01', 180), (2, 'K01', 190), (3, 'K02', 175);`} + // TRADITIONAL does not include ONLY_FULL_GROUP_BY, so this statement only + // succeeds when the init session's sql_mode reaches the graded statement. + result, err := server.executeQuery(context.Background(), init, "SELECT team, id, MAX(height) FROM kbo.PLAYER GROUP BY team ORDER BY team", nil) + if err != nil { + t.Fatal(err) + } + if len(result.Rows) != 2 || result.Rows[0][0] != "K01" || result.Rows[0][2] != "190" { + t.Fatalf("rows = %#v", result.Rows) + } + var count int + if err := server.controlDB.QueryRow("SELECT COUNT(*) FROM INFORMATION_SCHEMA.SCHEMATA WHERE SCHEMA_NAME = 'kbo'").Scan(&count); err != nil { + t.Fatal(err) + } + if count != 0 { + t.Fatal("init script created a real 'kbo' schema instead of using the temporary database") + } + carried, err := server.executeQuery(context.Background(), init, "SELECT @@SESSION.sql_mode", nil) + if err != nil { + t.Fatal(err) + } + if strings.Contains(carried.Rows[0][0], "ONLY_FULL_GROUP_BY") || !strings.Contains(carried.Rows[0][0], "STRICT_ALL_TABLES") { + t.Fatalf("graded session sql_mode = %q, want the init script's TRADITIONAL mode", carried.Rows[0][0]) + } + plain, err := server.executeQuery(context.Background(), nil, "SELECT @@SESSION.sql_mode", nil) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(plain.Rows[0][0], "ONLY_FULL_GROUP_BY") { + t.Fatalf("without an init script the server default sql_mode must apply, got %q", plain.Rows[0][0]) + } + assertNoTemporaryResources(t, server) +} + func TestIntegrationServerHardeningSettings(t *testing.T) { server := integrationServer(t) var eventScheduler, collation string @@ -189,6 +233,13 @@ func TestIntegrationServerHardeningSettings(t *testing.T) { if collation != temporaryCollation { t.Fatalf("collation_server = %q, want %q", collation, temporaryCollation) } + var lowerCaseTableNames int + if err := server.controlDB.QueryRow("SELECT @@lower_case_table_names").Scan(&lowerCaseTableNames); err != nil { + t.Fatal(err) + } + if lowerCaseTableNames != 1 { + t.Fatalf("lower_case_table_names = %d, want 1 (case-insensitive table names like Windows and macOS)", lowerCaseTableNames) + } } func TestIntegrationRejectsMultipleStatementsAndResultOverflow(t *testing.T) { diff --git a/platform/CTFd/plugins/sql_challenges/sql_judge_server.go b/platform/CTFd/plugins/sql_challenges/sql_judge_server.go index 38b031024e..8418169e68 100644 --- a/platform/CTFd/plugins/sql_challenges/sql_judge_server.go +++ b/platform/CTFd/plugins/sql_challenges/sql_judge_server.go @@ -33,7 +33,15 @@ var ( temporaryDatabasePattern = regexp.MustCompile(`^ctfd_tmp_[0-9a-f]{32}$`) temporaryUserPattern = regexp.MustCompile(`^ct_[0-9a-f]{16}$`) temporaryPasswordPattern = regexp.MustCompile(`^[0-9a-f]{64}$`) - errResultLimit = errors.New("query result exceeds configured limit") + // Challenge init scripts are usually exported from a local MySQL together + // with their own schema name. These statements are mapped onto the + // temporary database instead of being executed. + schemaStatementPattern = regexp.MustCompile("(?is)^(?:CREATE\\s+(?:DATABASE|SCHEMA)(?:\\s+IF\\s+NOT\\s+EXISTS)?|DROP\\s+(?:DATABASE|SCHEMA)(?:\\s+IF\\s+EXISTS)?|USE)\\s+`?([A-Za-z0-9_$]+)`?(?:\\s|$)") + sqlModePattern = regexp.MustCompile(`^[A-Za-z0-9_,]*$`) + // mysqldump wraps schema statements in version comments that MySQL executes, + // for example CREATE DATABASE /*!32312 IF NOT EXISTS*/ `world`. + versionedCommentPattern = regexp.MustCompile(`/\*!\d*\s*|\*/`) + errResultLimit = errors.New("query result exceeds configured limit") ) type QueryRequest struct { @@ -492,10 +500,11 @@ func (s *Server) executeQuery(parent context.Context, initQueries []string, quer if err := s.createExecutionResources(executionCtx, resources); err != nil { return nil, err } - if err := s.runInitStatements(executionCtx, resources, initQueries, req); err != nil { + session, err := s.runInitStatements(executionCtx, resources, initQueries, req) + if err != nil { return nil, err } - return s.runGradedQuery(executionCtx, resources, query) + return s.runGradedQuery(executionCtx, resources, query, session) } func (s *Server) openRunner(user, password, database string) (*sql.DB, error) { @@ -521,7 +530,17 @@ func (s *Server) openRunner(user, password, database string) (*sql.DB, error) { return runnerDB, nil } -func (s *Server) runInitStatements(ctx context.Context, r executionResources, initQueries []string, req *QueryRequest) error { +// initSession carries the parts of the init session that the graded statement +// must observe as if both ran in one session, which is how the previous +// in-process engine and a local MySQL client behave. +type initSession struct { + sqlMode string + sqlModeSet bool + schemaAliases []string +} + +func (s *Server) runInitStatements(ctx context.Context, r executionResources, initQueries []string, req *QueryRequest) (initSession, error) { + var session initSession var statements []string for _, initQuery := range initQueries { if strings.TrimSpace(initQuery) == "" { @@ -529,7 +548,7 @@ func (s *Server) runInitStatements(ctx context.Context, r executionResources, in } if err := validateSQLQuery(initQuery, req); err != nil { if strings.Contains(err.Error(), "file") || strings.Contains(err.Error(), "system") { - return fmt.Errorf("security violation in init query: %w", err) + return session, fmt.Errorf("security violation in init query: %w", err) } } for _, statement := range strings.Split(initQuery, ";") { @@ -539,23 +558,101 @@ func (s *Server) runInitStatements(ctx context.Context, r executionResources, in } } if len(statements) == 0 { - return nil + return session, nil + } + var executable []string + for _, statement := range statements { + // Leading comment lines are dropped before execution: MySQL treats a + // separator such as "-----" as a syntax error, and a chunk that is only + // comments would be an empty query. + statement = stripLeadingComments(statement) + if statement == "" { + continue + } + if name, ok := schemaStatementName(statement); ok { + session.schemaAliases = appendUnique(session.schemaAliases, name) + continue + } + executable = append(executable, statement) } initDB, err := s.openRunner(r.initUser, r.initPassword, r.database) if err != nil { - return err + return session, err } defer initDB.Close() - for _, statement := range statements { - if _, err := initDB.ExecContext(ctx, statement); err != nil { - return fmt.Errorf("init query error: %w", err) + // Pin one connection so SET statements in the script stay in effect. + conn, err := initDB.Conn(ctx) + if err != nil { + return session, fmt.Errorf("connect init MySQL user: %w", err) + } + defer conn.Close() + for _, statement := range executable { + statement = rewriteSchemaAliases(statement, session.schemaAliases, r.database) + if _, err := conn.ExecContext(ctx, statement); err != nil { + return session, fmt.Errorf("init query error: %w", err) } } - return nil + if err := conn.QueryRowContext(ctx, "SELECT @@SESSION.sql_mode").Scan(&session.sqlMode); err != nil { + return session, fmt.Errorf("read init session sql_mode: %w", err) + } + session.sqlModeSet = true + return session, nil +} + +// schemaStatementName reports the schema named by a CREATE/DROP DATABASE or +// USE statement after leading comments are removed. +func schemaStatementName(statement string) (string, bool) { + normalized := versionedCommentPattern.ReplaceAllString(stripLeadingComments(statement), " ") + match := schemaStatementPattern.FindStringSubmatch(strings.TrimSpace(normalized)) + if match == nil { + return "", false + } + return match[1], true +} + +func stripLeadingComments(statement string) string { + for { + statement = strings.TrimSpace(statement) + switch { + case strings.HasPrefix(statement, "--") || strings.HasPrefix(statement, "#"): + index := strings.IndexByte(statement, '\n') + if index < 0 { + return "" + } + statement = statement[index+1:] + case strings.HasPrefix(statement, "/*") && !strings.HasPrefix(statement, "/*!"): + index := strings.Index(statement, "*/") + if index < 0 { + return "" + } + statement = statement[index+2:] + default: + return statement + } + } +} + +// rewriteSchemaAliases points schema-qualified names from the init script at +// the temporary database, so `kbo.PLAYER` works in init and graded statements. +func rewriteSchemaAliases(statement string, aliases []string, database string) string { + for _, alias := range aliases { + pattern := regexp.MustCompile("(?i)(^|[^A-Za-z0-9_$.`])`?" + regexp.QuoteMeta(alias) + "`?\\s*\\.\\s*(`?[A-Za-z0-9_$]+`?)") + statement = pattern.ReplaceAllString(statement, "${1}"+quoteIdentifier(database)+".${2}") + } + return statement } -func (s *Server) runGradedQuery(ctx context.Context, r executionResources, query string) (*QueryResult, error) { +func appendUnique(values []string, value string) []string { + for _, existing := range values { + if strings.EqualFold(existing, value) { + return values + } + } + return append(values, value) +} + +func (s *Server) runGradedQuery(ctx context.Context, r executionResources, query string, session initSession) (*QueryResult, error) { queryDB, err := s.openRunner(r.queryUser, r.queryPassword, r.database) if err != nil { return nil, err @@ -580,6 +677,15 @@ func (s *Server) runGradedQuery(ctx context.Context, r executionResources, query if _, err := conn.ExecContext(ctx, fmt.Sprintf("SET SESSION max_execution_time = %d", maxExecutionMilliseconds)); err != nil { return nil, fmt.Errorf("set query execution limit: %w", err) } + if session.sqlModeSet { + if !sqlModePattern.MatchString(session.sqlMode) { + return nil, errors.New("init session left an unexpected sql_mode") + } + if _, err := conn.ExecContext(ctx, "SET SESSION sql_mode = '"+session.sqlMode+"'"); err != nil { + return nil, fmt.Errorf("apply init session sql_mode: %w", err) + } + } + query = rewriteSchemaAliases(query, session.schemaAliases, r.database) rows, err := conn.QueryContext(ctx, query) if err != nil { diff --git a/platform/CTFd/plugins/sql_challenges/sql_judge_server_test.go b/platform/CTFd/plugins/sql_challenges/sql_judge_server_test.go index bef41886f7..3a63f8bc38 100644 --- a/platform/CTFd/plugins/sql_challenges/sql_judge_server_test.go +++ b/platform/CTFd/plugins/sql_challenges/sql_judge_server_test.go @@ -163,3 +163,49 @@ func TestTemporaryNamesFitMySQLLimits(t *testing.T) { t.Fatal("temporary name quoting mixed SQL literal and identifier delimiters") } } + +func TestInitSchemaStatementsMapOntoTemporaryDatabase(t *testing.T) { + for statement, want := range map[string]string{ + "CREATE SCHEMA kbo": "kbo", + "create database if not exists `world` default character set utf8": "world", + "DROP SCHEMA IF EXISTS kbo": "kbo", + "USE kbo": "kbo", + "/* @@SESSION note */\nDROP SCHEMA IF EXISTS company": "company", + "-- KBO Database Schema\n-- Version 1.0\nUSE airport": "airport", + "-------------\n-- Schema\n-------------\nUSE company": "company", + "CREATE DATABASE /*!32312 IF NOT EXISTS*/ `world` /*!40100 DEFAULT CHARACTER SET utf8mb4 */": "world", + "--\n-- Current Database: `world`\n--\n/*!40000 DROP DATABASE IF EXISTS `world`*/": "world", + } { + got, ok := schemaStatementName(statement) + if !ok || got != want { + t.Fatalf("schemaStatementName(%q) = %q, %v; want %q", statement, got, ok, want) + } + } + for _, statement := range []string{ + "-- comment only", + "SET @OLD_SQL_MODE=@@SQL_MODE, SQL_MODE='TRADITIONAL'", + "CREATE TABLE kbo (id INT)", + "/*!40101 SET NAMES utf8 */", + "INSERT INTO t VALUES ('USE kbo')", + } { + if name, ok := schemaStatementName(statement); ok { + t.Fatalf("schemaStatementName(%q) unexpectedly matched %q", statement, name) + } + } + rewritten := rewriteSchemaAliases("SELECT * FROM kbo.PLAYER p JOIN `kbo`.`TEAM` t ON t.id = p.team_id WHERE akbo.x = 1", []string{"kbo"}, "ctfd_tmp_x") + if !strings.Contains(rewritten, "FROM `ctfd_tmp_x`.PLAYER") || !strings.Contains(rewritten, "JOIN `ctfd_tmp_x`.`TEAM`") || !strings.Contains(rewritten, "akbo.x") { + t.Fatalf("unexpected rewrite: %s", rewritten) + } + if got := rewriteSchemaAliases("SELECT 1", nil, "ctfd_tmp_x"); got != "SELECT 1" { + t.Fatalf("rewrite without aliases changed the statement: %s", got) + } + if got := stripLeadingComments("-----------\n-- Data\n-----------\nINSERT INTO t VALUES (1)"); got != "INSERT INTO t VALUES (1)" { + t.Fatalf("stripLeadingComments() = %q", got) + } + if got := stripLeadingComments("/*!40101 SET NAMES utf8 */"); !strings.HasPrefix(got, "/*!40101") { + t.Fatalf("version comment must be kept for execution, got %q", got) + } + if got := stripLeadingComments("-- only a comment"); got != "" { + t.Fatalf("comment-only chunk should be empty, got %q", got) + } +} diff --git a/platform/docker-compose.production.yml b/platform/docker-compose.production.yml index 48e691d7de..a6404c9c6d 100644 --- a/platform/docker-compose.production.yml +++ b/platform/docker-compose.production.yml @@ -48,6 +48,7 @@ services: - --collation-server=utf8mb4_0900_ai_ci - --disable-log-bin - --event-scheduler=DISABLED + - --lower-case-table-names=1 - --innodb-buffer-pool-size=128M - --innodb-flush-log-at-trx-commit=2 - --max-allowed-packet=16M diff --git a/platform/docker-compose.yml b/platform/docker-compose.yml index b16c39b081..c4f89e035c 100644 --- a/platform/docker-compose.yml +++ b/platform/docker-compose.yml @@ -74,6 +74,7 @@ services: - --collation-server=utf8mb4_0900_ai_ci - --disable-log-bin - --event-scheduler=DISABLED + - --lower-case-table-names=1 - --innodb-buffer-pool-size=128M - --innodb-flush-log-at-trx-commit=2 - --max-allowed-packet=16M diff --git a/scripts/regrade-challenges b/scripts/regrade-challenges index 0271a0d122..cc30bdc72b 100755 --- a/scripts/regrade-challenges +++ b/scripts/regrade-challenges @@ -24,6 +24,7 @@ from __future__ import annotations import argparse import concurrent.futures import json +import math import sys import time import urllib.error @@ -129,8 +130,9 @@ def classify_error(message: str) -> str: def numeric_equal(left: str, right: str) -> bool: + """Equal as numbers, allowing for DECIMAL vs float text precision (72220.1111 vs 72220.11111111111).""" try: - return float(left) == float(right) + return math.isclose(float(left), float(right), rel_tol=1e-6, abs_tol=1e-9) except ValueError: return False diff --git a/tests/test_sql_judge_runtime.py b/tests/test_sql_judge_runtime.py index 734a14c08f..d9eb7e8af1 100644 --- a/tests/test_sql_judge_runtime.py +++ b/tests/test_sql_judge_runtime.py @@ -35,6 +35,7 @@ def test_mysql_is_pinned_private_and_resource_limited(self): self.assertIn("judge-db", mysql) self.assertIn("--disable-log-bin", mysql) self.assertIn("--event-scheduler=DISABLED", mysql) + self.assertIn("--lower-case-table-names=1", mysql) self.assertIn("--collation-server=utf8mb4_0900_ai_ci", mysql) self.assertIn("--innodb-flush-log-at-trx-commit=2", mysql) self.assertIn("--max-connections=64", mysql)