修复: 升级框架并完善报告导出

- 升级 Drogon 和 Trantor,修复畸形请求导致的连接计数泄漏\n- 增加第三方框架版本校验与自动重建\n- 完善完整报告导出和接口文档
This commit is contained in:
cloud
2026-08-10 09:50:09 +08:00
parent 99ed321d24
commit 0e28826073
82 changed files with 3095 additions and 566 deletions
@@ -78,6 +78,7 @@ struct CalculationOptions {
double power_step = 5.0; double power_step = 5.0;
double cleaning_wind_speed_step = 0.25; double cleaning_wind_speed_step = 0.25;
double curve_wind_speed_step = 0.5; double curve_wind_speed_step = 0.5;
double report_wind_speed_interval = 0.25;
double wind_speed_change_threshold = 1.0; double wind_speed_change_threshold = 1.0;
double iqr_lower_multiplier = 1.2; double iqr_lower_multiplier = 1.2;
double iqr_upper_multiplier = 2.0; double iqr_upper_multiplier = 2.0;
@@ -119,9 +120,11 @@ struct SchemeInfo {
double scheme_one_generator_speed_k = 0.9; double scheme_one_generator_speed_k = 0.9;
double scheme_one_rotor_radius = 78.0; double scheme_one_rotor_radius = 78.0;
double scheme_one_gearbox_ratio = 162.0; double scheme_one_gearbox_ratio = 162.0;
double scheme_one_report_wind_speed_interval = 0.25;
double grid_connected_speed = 1030.0; double grid_connected_speed = 1030.0;
double rated_generator_speed = 1755.0; double rated_generator_speed = 1755.0;
double rated_power = 2000.0; double rated_power = 2000.0;
double scheme_two_report_wind_speed_interval = 0.25;
}; };
constexpr double kRatedCornerWindBefore = 0.5; constexpr double kRatedCornerWindBefore = 0.5;
@@ -527,11 +530,14 @@ json SchemeToJson(const SchemeInfo& scheme) {
{"generator_speed_k", scheme.scheme_one_generator_speed_k}, {"generator_speed_k", scheme.scheme_one_generator_speed_k},
{"rotor_radius", scheme.scheme_one_rotor_radius}, {"rotor_radius", scheme.scheme_one_rotor_radius},
{"gearbox_ratio", scheme.scheme_one_gearbox_ratio}, {"gearbox_ratio", scheme.scheme_one_gearbox_ratio},
{"report_wind_speed_interval", scheme.scheme_one_report_wind_speed_interval},
}; };
} else if (scheme.id == kSchemeTwoId) { } else if (scheme.id == kSchemeTwoId) {
data["parameters"]["grid_connected_speed"] = scheme.grid_connected_speed; data["parameters"]["grid_connected_speed"] = scheme.grid_connected_speed;
data["parameters"]["rated_generator_speed"] = scheme.rated_generator_speed; data["parameters"]["rated_generator_speed"] = scheme.rated_generator_speed;
data["parameters"]["rated_power"] = scheme.rated_power; data["parameters"]["rated_power"] = scheme.rated_power;
data["parameters"]["report_wind_speed_interval"] =
scheme.scheme_two_report_wind_speed_interval;
} }
return data; return data;
} }
@@ -590,6 +596,8 @@ std::vector<SchemeInfo> LoadSchemes() {
load_non_negative("generator_speed_k", scheme.scheme_one_generator_speed_k); load_non_negative("generator_speed_k", scheme.scheme_one_generator_speed_k);
load_positive("rotor_radius", scheme.scheme_one_rotor_radius); load_positive("rotor_radius", scheme.scheme_one_rotor_radius);
load_positive("gearbox_ratio", scheme.scheme_one_gearbox_ratio); load_positive("gearbox_ratio", scheme.scheme_one_gearbox_ratio);
load_positive("report_wind_speed_interval",
scheme.scheme_one_report_wind_speed_interval);
} else if (scheme.id == kSchemeTwoId) { } else if (scheme.id == kSchemeTwoId) {
if (const auto value = GetNumberField(params, "grid_connected_speed"); if (const auto value = GetNumberField(params, "grid_connected_speed");
value.has_value() && value.value() > 0.0) { value.has_value() && value.value() > 0.0) {
@@ -603,6 +611,10 @@ std::vector<SchemeInfo> LoadSchemes() {
value.has_value() && value.value() > 0.0) { value.has_value() && value.value() > 0.0) {
scheme.rated_power = value.value(); scheme.rated_power = value.value();
} }
if (const auto value = GetNumberField(params, "report_wind_speed_interval");
value.has_value() && value.value() > 0.0 && value.value() <= 2.0) {
scheme.scheme_two_report_wind_speed_interval = value.value();
}
} }
} }
} }
@@ -790,6 +802,10 @@ CalculationOptions ParseOptions(const json& body) {
value.has_value() && value.value() > 0.0 && value.value() <= 2.0) { value.has_value() && value.value() > 0.0 && value.value() <= 2.0) {
options.curve_wind_speed_step = value.value(); options.curve_wind_speed_step = value.value();
} }
if (const auto value = GetDoubleField(opt, "report_wind_speed_interval");
value.has_value() && value.value() > 0.0 && value.value() <= 2.0) {
options.report_wind_speed_interval = value.value();
}
if (const auto value = GetDoubleField(opt, "wind_speed_change_threshold"); if (const auto value = GetDoubleField(opt, "wind_speed_change_threshold");
value.has_value() && value.value() >= 0.0) { value.has_value() && value.value() >= 0.0) {
options.wind_speed_change_threshold = value.value(); options.wind_speed_change_threshold = value.value();
@@ -1359,6 +1375,7 @@ void WindPowerController::SaveSchemeDescription(
std::optional<double> grid_connected_speed; std::optional<double> grid_connected_speed;
std::optional<double> rated_generator_speed; std::optional<double> rated_generator_speed;
std::optional<double> rated_power; std::optional<double> rated_power;
std::optional<double> report_wind_speed_interval;
if (normalized_id == kDefaultSchemeId) { if (normalized_id == kDefaultSchemeId) {
if (!body.value().contains("parameters") || !body.value()["parameters"].is_object()) { if (!body.value().contains("parameters") || !body.value()["parameters"].is_object()) {
SendError(callback, kErrorInvalidRequest, "方案一参数格式错误"); SendError(callback, kErrorInvalidRequest, "方案一参数格式错误");
@@ -1376,6 +1393,7 @@ void WindPowerController::SaveSchemeDescription(
if (!valid_positive("rated_power") || !valid_positive("rated_wind_speed") || if (!valid_positive("rated_power") || !valid_positive("rated_wind_speed") ||
!valid_positive("power_step") || !valid_positive("cleaning_wind_speed_step") || !valid_positive("power_step") || !valid_positive("cleaning_wind_speed_step") ||
!valid_positive("rotor_radius") || !valid_positive("gearbox_ratio") || !valid_positive("rotor_radius") || !valid_positive("gearbox_ratio") ||
!valid_positive("report_wind_speed_interval") ||
!valid_non_negative("wind_speed_change_threshold") || !valid_non_negative("wind_speed_change_threshold") ||
!valid_non_negative("iqr_lower_multiplier") || !valid_non_negative("iqr_lower_multiplier") ||
!valid_non_negative("iqr_upper_multiplier") || !valid_non_negative("iqr_upper_multiplier") ||
@@ -1385,6 +1403,7 @@ void WindPowerController::SaveSchemeDescription(
return; return;
} }
if (GetDoubleField(params, "cleaning_wind_speed_step").value() > 2.0 || if (GetDoubleField(params, "cleaning_wind_speed_step").value() > 2.0 ||
GetDoubleField(params, "report_wind_speed_interval").value() > 2.0 ||
GetDoubleField(params, "iqr_lower_multiplier").value() > 10.0 || GetDoubleField(params, "iqr_lower_multiplier").value() > 10.0 ||
GetDoubleField(params, "iqr_upper_multiplier").value() > 10.0) { GetDoubleField(params, "iqr_upper_multiplier").value() > 10.0) {
SendError(callback, kErrorInvalidRequest, "方案一参数超出允许范围"); SendError(callback, kErrorInvalidRequest, "方案一参数超出允许范围");
@@ -1400,9 +1419,12 @@ void WindPowerController::SaveSchemeDescription(
grid_connected_speed = GetDoubleField(params, "grid_connected_speed"); grid_connected_speed = GetDoubleField(params, "grid_connected_speed");
rated_generator_speed = GetDoubleField(params, "rated_generator_speed"); rated_generator_speed = GetDoubleField(params, "rated_generator_speed");
rated_power = GetDoubleField(params, "rated_power"); rated_power = GetDoubleField(params, "rated_power");
report_wind_speed_interval = GetDoubleField(params, "report_wind_speed_interval");
if (!grid_connected_speed.has_value() || grid_connected_speed.value() <= 0.0 || if (!grid_connected_speed.has_value() || grid_connected_speed.value() <= 0.0 ||
!rated_generator_speed.has_value() || rated_generator_speed.value() <= 0.0 || !rated_generator_speed.has_value() || rated_generator_speed.value() <= 0.0 ||
!rated_power.has_value() || rated_power.value() <= 0.0) { !rated_power.has_value() || rated_power.value() <= 0.0 ||
!report_wind_speed_interval.has_value() || report_wind_speed_interval.value() <= 0.0 ||
report_wind_speed_interval.value() > 2.0) {
SendError(callback, kErrorInvalidRequest, "方案二参数必须为正数"); SendError(callback, kErrorInvalidRequest, "方案二参数必须为正数");
return; return;
} }
@@ -1432,6 +1454,8 @@ void WindPowerController::SaveSchemeDescription(
GetDoubleField(params, "generator_speed_k").value(); GetDoubleField(params, "generator_speed_k").value();
scheme.scheme_one_rotor_radius = GetDoubleField(params, "rotor_radius").value(); scheme.scheme_one_rotor_radius = GetDoubleField(params, "rotor_radius").value();
scheme.scheme_one_gearbox_ratio = GetDoubleField(params, "gearbox_ratio").value(); scheme.scheme_one_gearbox_ratio = GetDoubleField(params, "gearbox_ratio").value();
scheme.scheme_one_report_wind_speed_interval =
GetDoubleField(params, "report_wind_speed_interval").value();
} else if (scheme.id == kSchemeTwoId) { } else if (scheme.id == kSchemeTwoId) {
if (grid_connected_speed.has_value()) { if (grid_connected_speed.has_value()) {
scheme.grid_connected_speed = grid_connected_speed.value(); scheme.grid_connected_speed = grid_connected_speed.value();
@@ -1442,6 +1466,10 @@ void WindPowerController::SaveSchemeDescription(
if (rated_power.has_value()) { if (rated_power.has_value()) {
scheme.rated_power = rated_power.value(); scheme.rated_power = rated_power.value();
} }
if (report_wind_speed_interval.has_value()) {
scheme.scheme_two_report_wind_speed_interval =
report_wind_speed_interval.value();
}
} }
updated = true; updated = true;
break; break;
@@ -1619,6 +1647,16 @@ void WindPowerController::FinishJob(
TaskReleaseGuard finish_guard(job_id.value()); TaskReleaseGuard finish_guard(job_id.value());
const CalculationOptions options = ParseOptions(*body); const CalculationOptions options = ParseOptions(*body);
if (body->contains("options") && (*body)["options"].is_object() &&
(*body)["options"].contains("report_wind_speed_interval")) {
const auto interval = GetDoubleField((*body)["options"], "report_wind_speed_interval");
if (!interval.has_value() || !std::isfinite(interval.value()) ||
interval.value() <= 0.0 || interval.value() > 2.0) {
SendError(callback, kErrorInvalidRequest,
"报告公式风速区间半宽必须大于 0 且不超过 2");
return;
}
}
if (IsSchemeTwo(options) && if (IsSchemeTwo(options) &&
(!options.rated_power_provided || (!options.rated_power_provided ||
!options.grid_connected_speed_provided || !options.grid_connected_speed_provided ||
@@ -1964,6 +2002,8 @@ void WindPowerController::FinishJob(
json data; json data;
const auto selected_scheme = FindScheme(LoadSchemes(), options.scheme_id); const auto selected_scheme = FindScheme(LoadSchemes(), options.scheme_id);
data["scheme"] = SchemeToJson(selected_scheme.value()); data["scheme"] = SchemeToJson(selected_scheme.value());
data["scheme"]["parameters"]["report_wind_speed_interval"] =
options.report_wind_speed_interval;
data["summary"] = summary; data["summary"] = summary;
data["fans"] = fans; data["fans"] = fans;
data["curves"] = curves; data["curves"] = curves;
@@ -2108,25 +2148,36 @@ void WindPowerController::ExportReport(
worksheet_set_column(curve, i, i, i == 0 ? 9 : 17, nullptr); worksheet_set_column(curve, i, i, i == 0 ? 9 : 17, nullptr);
} }
const auto raw_wind = ExcelColumnName(wind_column); const auto raw_wind = ExcelColumnName(wind_column);
const auto report_interval = result.value("scheme", json::object())
.value("parameters", json::object())
.value("report_wind_speed_interval", 0.25);
const auto interval_text = std::to_string(report_interval);
const auto report_row_count = (*body)["report_rows"].size();
const auto k_last_row = std::max<size_t>(2, report_row_count + 1);
const auto k_formula = "=IFERROR(SUM(F2:F" + std::to_string(k_last_row) +
")/SUM(G2:G" + std::to_string(k_last_row) + "),0)";
lxw_row_t curve_row = 1; lxw_row_t curve_row = 1;
for (const auto& point : (*body)["report_rows"]) { for (const auto& point : (*body)["report_rows"]) {
const auto excel_row = curve_row + 1; const auto excel_row = curve_row + 1;
worksheet_write_number(curve, curve_row, 0, excel_row - 1, nullptr); worksheet_write_number(curve, curve_row, 0, excel_row - 1, nullptr);
worksheet_write_number(curve, curve_row, 1, point.value("wind_speed", 0.0), number_format); worksheet_write_number(curve, curve_row, 1, point.value("wind_speed", 0.0), number_format);
const auto frequency = "=(COUNTIFS('筛选前的数据'!" + raw_wind + ":" + raw_wind + ",\">=\"&B" + std::to_string(excel_row) + "-0.5,'筛选前的数据'!" + raw_wind + ":" + raw_wind + ",\"<\"&B" + std::to_string(excel_row) + "+0.5)/COUNT('筛选前的数据'!" + raw_wind + ":" + raw_wind + "))*8760"; const auto frequency = "=(COUNTIFS('筛选前的数据'!" + raw_wind + ":" + raw_wind + ",\">=\"&B" + std::to_string(excel_row) + "-" + interval_text + ",'筛选前的数据'!" + raw_wind + ":" + raw_wind + ",\"<\"&B" + std::to_string(excel_row) + "+" + interval_text + ")/COUNT('筛选前的数据'!" + raw_wind + ":" + raw_wind + "))*8760";
worksheet_write_formula(curve, curve_row, 2, frequency.c_str(), number_format); worksheet_write_formula(curve, curve_row, 2, frequency.c_str(), number_format);
const auto actual = "=IFERROR(AVERAGEIFS('筛选后的数据'!$C:$C,'筛选后的数据'!$E:$E,\">=\"&B" + std::to_string(excel_row) + "-0.5,'筛选后的数据'!$E:$E,\"<\"&B" + std::to_string(excel_row) + "+0.5),0)"; const auto actual = "=IFERROR(AVERAGEIFS('筛选后的数据'!$C:$C,'筛选后的数据'!$E:$E,\">=\"&B" + std::to_string(excel_row) + "-" + interval_text + ",'筛选后的数据'!$E:$E,\"<\"&B" + std::to_string(excel_row) + "+" + interval_text + "),0)";
worksheet_write_formula(curve, curve_row, 3, actual.c_str(), number_format); worksheet_write_formula(curve, curve_row, 3, actual.c_str(), number_format);
worksheet_write_number(curve, curve_row, 4, point.value("design_power", 0.0), number_format); worksheet_write_number(curve, curve_row, 4, point.value("design_power", 0.0), number_format);
const auto generated = "=IFERROR(ROUND(C" + std::to_string(excel_row) + "*D" + std::to_string(excel_row) + "/1000,4),0)"; const auto generated = "=IFERROR(ROUND(C" + std::to_string(excel_row) + "*D" + std::to_string(excel_row) + "/1000,4),0)";
const auto theoretical = "=ROUND(C" + std::to_string(excel_row) + "*E" + std::to_string(excel_row) + "/1000,4)"; const auto theoretical = "=ROUND(C" + std::to_string(excel_row) + "*E" + std::to_string(excel_row) + "/1000,4)";
worksheet_write_formula(curve, curve_row, 5, generated.c_str(), number_format); worksheet_write_formula(curve, curve_row, 5, generated.c_str(), number_format);
worksheet_write_formula(curve, curve_row, 6, theoretical.c_str(), number_format); worksheet_write_formula(curve, curve_row, 6, theoretical.c_str(), number_format);
if (curve_row == 1) {
worksheet_write_formula(curve, curve_row, 8, k_formula.c_str(), number_format);
}
++curve_row; ++curve_row;
} }
const auto last_row = std::max<lxw_row_t>(2, curve_row); if (report_row_count == 0) {
const auto k_formula = "=IFERROR(SUM(F2:F" + std::to_string(last_row) + ")/SUM(G2:G" + std::to_string(last_row) + "),0)"; worksheet_write_formula(curve, 1, 8, k_formula.c_str(), number_format);
worksheet_write_formula(curve, 1, 8, k_formula.c_str(), number_format); }
worksheet_autofilter(curve, 0, 0, std::max<lxw_row_t>(1, curve_row - 1), 6); worksheet_autofilter(curve, 0, 0, std::max<lxw_row_t>(1, curve_row - 1), 6);
if (body->contains("chart_image") && (*body)["chart_image"].is_string()) { if (body->contains("chart_image") && (*body)["chart_image"].is_string()) {
@@ -2137,7 +2188,11 @@ void WindPowerController::ExportReport(
std::ofstream image(image_path, std::ios::binary | std::ios::trunc); std::ofstream image(image_path, std::ios::binary | std::ios::trunc);
image << drogon::utils::base64Decode(image_data); image << drogon::utils::base64Decode(image_data);
image.close(); image.close();
worksheet_insert_image(curve, 3, 8, image_path.string().c_str()); lxw_image_options image_options{};
image_options.x_scale = 0.5;
image_options.y_scale = 0.5;
worksheet_insert_image_opt(curve, 3, 8, image_path.string().c_str(),
&image_options);
} }
if (workbook_close(workbook) != LXW_NO_ERROR) throw std::runtime_error("写入 Excel 文件失败"); if (workbook_close(workbook) != LXW_NO_ERROR) throw std::runtime_error("写入 Excel 文件失败");
callback(HttpResponse::newFileResponse(report_path.string(), callback(HttpResponse::newFileResponse(report_path.string(),
+3 -1
View File
@@ -146,7 +146,8 @@
"minimum_generator_speed": 1, "minimum_generator_speed": 1,
"generator_speed_k": 0.9, "generator_speed_k": 0.9,
"rotor_radius": 78, "rotor_radius": 78,
"gearbox_ratio": 162 "gearbox_ratio": 162,
"report_wind_speed_interval": 0.25
} }
} }
``` ```
@@ -244,6 +245,7 @@
- 基于分箱中位功率曲线的残差异常点剔除为 `curve_residual_outlier` - 基于分箱中位功率曲线的残差异常点剔除为 `curve_residual_outlier`
- 严格额定平台区残留偏低点剔除为 `rated_plateau_low_power` - 严格额定平台区残留偏低点剔除为 `rated_plateau_low_power`
- 最终曲线按 `curve_wind_speed_step` 左开右闭分箱,区间非空即输出平均功率点。 - 最终曲线按 `curve_wind_speed_step` 左开右闭分箱,区间非空即输出平均功率点。
- `report_wind_speed_interval` 为完整报告 Sheet3 的公式区间半宽,默认 `0.25 m/s`,取值范围 `(0, 2]`;方案一、方案二分别保存。
- `scatter_points` 返回清洗后保留点,`filtered_points` 返回所有过滤阶段滤除的点和原因。 - `scatter_points` 返回清洗后保留点,`filtered_points` 返回所有过滤阶段滤除的点和原因。
- `estimated_params.source``auto``auto_power_fallback_wind``fallback` - `estimated_params.source``auto``auto_power_fallback_wind``fallback`
+40 -3
View File
@@ -75,6 +75,7 @@ const DEFAULT_SCHEMES = [
generator_speed_k: 0.9, generator_speed_k: 0.9,
rotor_radius: 78, rotor_radius: 78,
gearbox_ratio: 162, gearbox_ratio: 162,
report_wind_speed_interval: 0.25,
}, },
}, },
{ {
@@ -85,6 +86,7 @@ const DEFAULT_SCHEMES = [
grid_connected_speed: 1030, grid_connected_speed: 1030,
rated_generator_speed: 1755, rated_generator_speed: 1755,
rated_power: 2000, rated_power: 2000,
report_wind_speed_interval: 0.25,
}, },
}, },
]; ];
@@ -93,6 +95,7 @@ const DEFAULT_SCHEME_TWO_PARAMS = {
grid_connected_speed: '1030', grid_connected_speed: '1030',
rated_generator_speed: '1755', rated_generator_speed: '1755',
rated_power: '2000', rated_power: '2000',
report_wind_speed_interval: '0.25',
}; };
const DEFAULT_SCHEME_ONE_PARAMS = { const DEFAULT_SCHEME_ONE_PARAMS = {
rated_power: '4800', rated_power: '4800',
@@ -106,6 +109,7 @@ const DEFAULT_SCHEME_ONE_PARAMS = {
generator_speed_k: '0.9', generator_speed_k: '0.9',
rotor_radius: '78', rotor_radius: '78',
gearbox_ratio: '162', gearbox_ratio: '162',
report_wind_speed_interval: '0.25',
}; };
const SCHEME_ONE_PARAM_FIELDS = [ const SCHEME_ONE_PARAM_FIELDS = [
{ key: 'rated_power', label: '额定功率', unit: 'kW', min: '0', step: '0.1' }, { key: 'rated_power', label: '额定功率', unit: 'kW', min: '0', step: '0.1' },
@@ -119,6 +123,7 @@ const SCHEME_ONE_PARAM_FIELDS = [
{ key: 'generator_speed_k', label: '发电机转速系数 K', unit: '', min: '0', step: '0.01' }, { key: 'generator_speed_k', label: '发电机转速系数 K', unit: '', min: '0', step: '0.01' },
{ key: 'rotor_radius', label: '叶轮半径', unit: 'm', min: '0', step: '0.01' }, { key: 'rotor_radius', label: '叶轮半径', unit: 'm', min: '0', step: '0.01' },
{ key: 'gearbox_ratio', label: '齿轮箱传动比', unit: '', min: '0', step: '0.01' }, { key: 'gearbox_ratio', label: '齿轮箱传动比', unit: '', min: '0', step: '0.01' },
{ key: 'report_wind_speed_interval', label: '报告公式风速区间半宽', unit: 'm/s', min: '0', step: '0.01' },
]; ];
const CHART_STATE_KEY = 'wind_power_chart_state_v1'; const CHART_STATE_KEY = 'wind_power_chart_state_v1';
const DEFAULT_CHART_OPTIONS = { const DEFAULT_CHART_OPTIONS = {
@@ -2045,6 +2050,8 @@ export default function HomePage() {
chart_image: chartImageData || '', chart_image: chartImageData || '',
}); });
downloadReportBlob(blob, selectedFan); downloadReportBlob(blob, selectedFan);
setExportingReport(false);
setReportExportProgress('');
} catch (err) { } catch (err) {
const message = err.message || '导出完整报告失败'; const message = err.message || '导出完整报告失败';
setError(message); setError(message);
@@ -2141,6 +2148,7 @@ export default function HomePage() {
const gridConnectedSpeed = normalizeNumber(schemeTwoParams.grid_connected_speed); const gridConnectedSpeed = normalizeNumber(schemeTwoParams.grid_connected_speed);
const ratedGeneratorSpeed = normalizeNumber(schemeTwoParams.rated_generator_speed); const ratedGeneratorSpeed = normalizeNumber(schemeTwoParams.rated_generator_speed);
const ratedPower = normalizeNumber(schemeTwoParams.rated_power); const ratedPower = normalizeNumber(schemeTwoParams.rated_power);
const reportWindSpeedInterval = normalizeNumber(schemeTwoParams.report_wind_speed_interval);
if (selectedScheme.id === DEFAULT_SCHEME_ID && if (selectedScheme.id === DEFAULT_SCHEME_ID &&
(!Number.isFinite(schemeOneValues.rated_power) || schemeOneValues.rated_power <= 0 || (!Number.isFinite(schemeOneValues.rated_power) || schemeOneValues.rated_power <= 0 ||
!Number.isFinite(schemeOneValues.rated_wind_speed) || schemeOneValues.rated_wind_speed <= 0 || !Number.isFinite(schemeOneValues.rated_wind_speed) || schemeOneValues.rated_wind_speed <= 0 ||
@@ -2148,6 +2156,9 @@ export default function HomePage() {
!Number.isFinite(schemeOneValues.cleaning_wind_speed_step) || schemeOneValues.cleaning_wind_speed_step <= 0 || !Number.isFinite(schemeOneValues.cleaning_wind_speed_step) || schemeOneValues.cleaning_wind_speed_step <= 0 ||
!Number.isFinite(schemeOneValues.rotor_radius) || schemeOneValues.rotor_radius <= 0 || !Number.isFinite(schemeOneValues.rotor_radius) || schemeOneValues.rotor_radius <= 0 ||
!Number.isFinite(schemeOneValues.gearbox_ratio) || schemeOneValues.gearbox_ratio <= 0 || !Number.isFinite(schemeOneValues.gearbox_ratio) || schemeOneValues.gearbox_ratio <= 0 ||
!Number.isFinite(schemeOneValues.report_wind_speed_interval) ||
schemeOneValues.report_wind_speed_interval <= 0 ||
schemeOneValues.report_wind_speed_interval > 2 ||
['wind_speed_change_threshold', 'iqr_lower_multiplier', 'iqr_upper_multiplier', ['wind_speed_change_threshold', 'iqr_lower_multiplier', 'iqr_upper_multiplier',
'minimum_generator_speed', 'generator_speed_k'].some((field) => ( 'minimum_generator_speed', 'generator_speed_k'].some((field) => (
!Number.isFinite(schemeOneValues[field]) || schemeOneValues[field] < 0 !Number.isFinite(schemeOneValues[field]) || schemeOneValues[field] < 0
@@ -2155,9 +2166,11 @@ export default function HomePage() {
throw new Error('方案一参数必须为合法数值,额定与步长参数、叶轮半径和传动比必须大于 0'); throw new Error('方案一参数必须为合法数值,额定与步长参数、叶轮半径和传动比必须大于 0');
} }
if (selectedScheme.id === SCHEME_TWO_ID && if (selectedScheme.id === SCHEME_TWO_ID &&
(!Number.isFinite(gridConnectedSpeed) || gridConnectedSpeed <= 0 || (!Number.isFinite(gridConnectedSpeed) || gridConnectedSpeed <= 0 ||
!Number.isFinite(ratedGeneratorSpeed) || ratedGeneratorSpeed <= 0 || !Number.isFinite(ratedGeneratorSpeed) || ratedGeneratorSpeed <= 0 ||
!Number.isFinite(ratedPower) || ratedPower <= 0)) { !Number.isFinite(ratedPower) || ratedPower <= 0 ||
!Number.isFinite(reportWindSpeedInterval) || reportWindSpeedInterval <= 0 ||
reportWindSpeedInterval > 2)) {
throw new Error('方案二参数必须为正数'); throw new Error('方案二参数必须为正数');
} }
const data = await saveWindSchemeDescription(selectedScheme.id, { const data = await saveWindSchemeDescription(selectedScheme.id, {
@@ -2168,6 +2181,7 @@ export default function HomePage() {
grid_connected_speed: gridConnectedSpeed, grid_connected_speed: gridConnectedSpeed,
rated_generator_speed: ratedGeneratorSpeed, rated_generator_speed: ratedGeneratorSpeed,
rated_power: ratedPower, rated_power: ratedPower,
report_wind_speed_interval: reportWindSpeedInterval,
}, },
} : {}), } : {}),
}); });
@@ -2288,6 +2302,9 @@ export default function HomePage() {
schemeOneValues.cleaning_wind_speed_step <= 0 || schemeOneValues.cleaning_wind_speed_step <= 0 ||
!Number.isFinite(schemeOneValues.rotor_radius) || schemeOneValues.rotor_radius <= 0 || !Number.isFinite(schemeOneValues.rotor_radius) || schemeOneValues.rotor_radius <= 0 ||
!Number.isFinite(schemeOneValues.gearbox_ratio) || schemeOneValues.gearbox_ratio <= 0 || !Number.isFinite(schemeOneValues.gearbox_ratio) || schemeOneValues.gearbox_ratio <= 0 ||
!Number.isFinite(schemeOneValues.report_wind_speed_interval) ||
schemeOneValues.report_wind_speed_interval <= 0 ||
schemeOneValues.report_wind_speed_interval > 2 ||
['wind_speed_change_threshold', 'iqr_lower_multiplier', 'iqr_upper_multiplier', ['wind_speed_change_threshold', 'iqr_lower_multiplier', 'iqr_upper_multiplier',
'minimum_generator_speed', 'generator_speed_k'].some((field) => ( 'minimum_generator_speed', 'generator_speed_k'].some((field) => (
!Number.isFinite(schemeOneValues[field]) || schemeOneValues[field] < 0 !Number.isFinite(schemeOneValues[field]) || schemeOneValues[field] < 0
@@ -2299,10 +2316,13 @@ export default function HomePage() {
const gridConnectedSpeed = normalizeNumber(schemeTwoParams.grid_connected_speed); const gridConnectedSpeed = normalizeNumber(schemeTwoParams.grid_connected_speed);
const ratedGeneratorSpeed = normalizeNumber(schemeTwoParams.rated_generator_speed); const ratedGeneratorSpeed = normalizeNumber(schemeTwoParams.rated_generator_speed);
const ratedPower = normalizeNumber(schemeTwoParams.rated_power); const ratedPower = normalizeNumber(schemeTwoParams.rated_power);
const reportWindSpeedInterval = normalizeNumber(schemeTwoParams.report_wind_speed_interval);
if (isSchemeTwo && if (isSchemeTwo &&
(!Number.isFinite(gridConnectedSpeed) || gridConnectedSpeed <= 0 || (!Number.isFinite(gridConnectedSpeed) || gridConnectedSpeed <= 0 ||
!Number.isFinite(ratedGeneratorSpeed) || ratedGeneratorSpeed <= 0 || !Number.isFinite(ratedGeneratorSpeed) || ratedGeneratorSpeed <= 0 ||
!Number.isFinite(ratedPower) || ratedPower <= 0)) { !Number.isFinite(ratedPower) || ratedPower <= 0 ||
!Number.isFinite(reportWindSpeedInterval) || reportWindSpeedInterval <= 0 ||
reportWindSpeedInterval > 2)) {
setError('方案二需要填写并网转速、额定转速和额定功率,且必须为正数'); setError('方案二需要填写并网转速、额定转速和额定功率,且必须为正数');
return; return;
} }
@@ -2354,6 +2374,7 @@ export default function HomePage() {
grid_connected_speed: gridConnectedSpeed, grid_connected_speed: gridConnectedSpeed,
rated_generator_speed: ratedGeneratorSpeed, rated_generator_speed: ratedGeneratorSpeed,
rated_power: ratedPower, rated_power: ratedPower,
report_wind_speed_interval: reportWindSpeedInterval,
} : {}), } : {}),
}, },
}); });
@@ -2520,6 +2541,22 @@ export default function HomePage() {
placeholder="请输入" placeholder="请输入"
/> />
</label> </label>
<label className="fieldControl">
<span>报告公式风速区间半宽 (m/s)</span>
<input
type="number"
min="0"
max="2"
step="0.01"
value={schemeTwoParams.report_wind_speed_interval}
onChange={(event) => handleSchemeTwoParamChange(
'report_wind_speed_interval',
event.target.value,
)}
disabled={submitting || savingScheme}
placeholder="请输入"
/>
</label>
</div> </div>
)} )}
<div className="schemeFooter"> <div className="schemeFooter">
+79 -6
View File
@@ -100,14 +100,44 @@ jobs:
- name: Prepare for testing - name: Prepare for testing
run: | run: |
brew services restart postgresql@14 brew services restart postgresql@14
for _ in {1..30}; do
if pg_isready -h 127.0.0.1 -p 5432 -U postgres; then
break
fi
sleep 1
done
pg_isready -h 127.0.0.1 -p 5432 -U postgres
brew services start mariadb brew services start mariadb
for _ in {1..30}; do
if mariadb-admin ping --silent; then
break
fi
sleep 1
done
mariadb-admin ping --silent
brew services start redis brew services start redis
sleep 4 for _ in {1..30}; do
if redis-cli ping | grep -q PONG; then
break
fi
sleep 1
done
redis-cli ping | grep -q PONG
mariadb -e "SET PASSWORD FOR 'root'@'localhost' = PASSWORD('')" mariadb -e "SET PASSWORD FOR 'root'@'localhost' = PASSWORD('')"
mariadb -e "GRANT ALL PRIVILEGES ON *.* TO 'root'@'localhost'" mariadb -e "GRANT ALL PRIVILEGES ON *.* TO 'root'@'localhost'"
mariadb -e "FLUSH PRIVILEGES" mariadb -e "FLUSH PRIVILEGES"
brew services restart mariadb brew services restart mariadb
sleep 4 for _ in {1..30}; do
if mariadb-admin ping --silent; then
break
fi
sleep 1
done
mariadb-admin ping --silent
psql -c 'create user postgres superuser;' postgres psql -c 'create user postgres superuser;' postgres
- name: Test - name: Test
@@ -178,9 +208,46 @@ jobs:
- name: Install g++ - name: Install g++
if: startsWith(matrix.compiler.cxx, 'g++') && (matrix.compiler.ver == 13 || matrix.compiler.ver == 9) if: startsWith(matrix.compiler.cxx, 'g++') && (matrix.compiler.ver == 13 || matrix.compiler.ver == 9)
run: | run: |
sudo add-apt-repository ppa:ubuntu-toolchain-r/test requested="g++-${{ matrix.compiler.ver }}"
sudo apt-get install g++-${{ matrix.compiler.ver }}
sudo update-alternatives --install /usr/bin/g++ g++ /usr/bin/g++-${{ matrix.compiler.ver }} ${{ matrix.compiler.ver }} if command -v g++-${{ matrix.compiler.ver }} >/dev/null 2>&1; then
echo "g++-${{ matrix.compiler.ver }} is already available on the runner"
else
retry() {
local attempts="$1"
shift
local try=1
while true; do
"$@" && break
if [ "$try" -ge "$attempts" ]; then
return 1
fi
echo "Command failed, retrying ($try/$attempts): $*"
try=$((try + 1))
sleep 10
done
}
if retry 3 sudo add-apt-repository -y ppa:ubuntu-toolchain-r/test && \
retry 3 sudo apt-get update && \
retry 3 sudo apt-get install -y "$requested"; then
echo "Installed $requested from PPA"
else
echo "::warning::Failed to install $requested from Launchpad PPA (network timeout). Falling back to system g++."
if ! command -v g++ >/dev/null 2>&1; then
echo "::error::Neither $requested nor system g++ is available."
exit 1
fi
fallback_cxx="$(command -v g++)"
echo "CXX=$fallback_cxx" >> "$GITHUB_ENV"
echo "Using fallback compiler: $fallback_cxx"
exit 0
fi
fi
if command -v "$requested" >/dev/null 2>&1; then
sudo update-alternatives --install /usr/bin/g++ g++ "$(command -v "$requested")" ${{ matrix.compiler.ver }}
fi
- name: Install Clang - name: Install Clang
if: startsWith(matrix.compiler.cxx, 'clang') && matrix.compiler.ver < 13 if: startsWith(matrix.compiler.cxx, 'clang') && matrix.compiler.ver < 13
@@ -228,7 +295,13 @@ jobs:
- name: Prepare for testing - name: Prepare for testing
run: | run: |
sudo systemctl start postgresql sudo systemctl start postgresql
sleep 1 for _ in {1..30}; do
if pg_isready -h 127.0.0.1 -p 5432 -U postgres; then
break
fi
sleep 1
done
pg_isready -h 127.0.0.1 -p 5432 -U postgres
sudo -u postgres psql -c "ALTER USER postgres WITH PASSWORD '12345'" postgres sudo -u postgres psql -c "ALTER USER postgres WITH PASSWORD '12345'" postgres
- name: Test - name: Test
+2
View File
@@ -47,3 +47,5 @@ CMakeSettings.json
install install
trace.json trace.json
.cache/ .cache/
build_examples/
.kiro
+4 -1
View File
@@ -25,7 +25,7 @@ CMAKE_DEPENDENT_OPTION(USE_SPDLOG "Allow using the spdlog logging library" OFF "
set(DROGON_MAJOR_VERSION 1) set(DROGON_MAJOR_VERSION 1)
set(DROGON_MINOR_VERSION 9) set(DROGON_MINOR_VERSION 9)
set(DROGON_PATCH_VERSION 12) set(DROGON_PATCH_VERSION 13)
set(DROGON_VERSION set(DROGON_VERSION
${DROGON_MAJOR_VERSION}.${DROGON_MINOR_VERSION}.${DROGON_PATCH_VERSION}) ${DROGON_MAJOR_VERSION}.${DROGON_MINOR_VERSION}.${DROGON_PATCH_VERSION})
set(DROGON_VERSION_STRING "${DROGON_VERSION}") set(DROGON_VERSION_STRING "${DROGON_VERSION}")
@@ -134,6 +134,8 @@ if (WIN32)
PRIVATE $<BUILD_INTERFACE:${PROJECT_SOURCE_DIR}/third_party/mman-win32>) PRIVATE $<BUILD_INTERFACE:${PROJECT_SOURCE_DIR}/third_party/mman-win32>)
endif (WIN32) endif (WIN32)
add_library(Drogon::Drogon ALIAS ${PROJECT_NAME})
if(USE_SUBMODULE) if(USE_SUBMODULE)
add_subdirectory(trantor) add_subdirectory(trantor)
target_link_libraries(${PROJECT_NAME} PUBLIC trantor) target_link_libraries(${PROJECT_NAME} PUBLIC trantor)
@@ -329,6 +331,7 @@ set(private_headers
lib/src/ListenerManager.h lib/src/ListenerManager.h
lib/src/PluginsManager.h lib/src/PluginsManager.h
lib/src/SessionManager.h lib/src/SessionManager.h
lib/src/utils/ParsingUtils.h
lib/src/SpinLock.h lib/src/SpinLock.h
lib/src/StaticFileRouter.h lib/src/StaticFileRouter.h
lib/src/TaskTimeoutFlag.h lib/src/TaskTimeoutFlag.h
+75 -1
View File
@@ -4,6 +4,78 @@ All notable changes to this project will be documented in this file.
## [Unreleased] ## [Unreleased]
## [1.9.13] - 2026-05-06
### API changes list
- Add `HttpRequest::clearHeaders()` method.
- Add `setQueryParameter()` and `setBodyParameter()` methods.
- Add ability to use `BEGIN IMMEDIATE` and `BEGIN EXCLUSIVE`.
- Add `UploadFile` constructor to create from memory data.
- Add JOIN support to Mapper/BaseBuilder with FK auto-detection in code generator.
- Add WebDAV HTTP methods (PROPFIND, MKCOL, COPY, MOVE).
- Add per-request compression control to `HttpResponse`.
### Added
- Support for custom OPTIONS handling via middleware flagging.
- Alias library for Drogon with name matching installed target.
- `--clear-output` option to drogon_ctl create models.
### Changed
- vector: reserve before inserting for efficiency.
- make `utils::isBase64` support padding.
- Extract duplicate `parseLine()` function to shared utility header.
- Doxygen documentation adjustments.
- Enhancement for custom OPTIONS handling.
- Forward the path methods.
### Fixed
- Fix sqlite3 test in CI.
- Fix HttpClient not sending WebDAV requests.
- Fix connection limit bug.
- Fix bugs exposed by CI.
- Fix HTTP date formatting to be locale-independent.
- Fix parsing invalid numbers in HTTP headers.
- Fix shared lib view failure handling.
- Fix drogon_ctl compilation with clang-cl.
- Fix missing throw statement.
- Add Homebrew Apple Silicon path detection in CMake finder modules.
- Include missing header files.
- Fix wrong numeric limit for floating types.
- Fix regex WebSocket routes middleware initialization.
- Fix system() replaced with execvp() in SharedLibManager.
- Fix inverted test logic.
## [1.9.12] - 2026-01-26 ## [1.9.12] - 2026-01-26
### API changes list ### API changes list
@@ -1878,7 +1950,9 @@ All notable changes to this project will be documented in this file.
## [1.0.0-beta1] - 2019-06-11 ## [1.0.0-beta1] - 2019-06-11
[Unreleased]: https://github.com/an-tao/drogon/compare/v1.9.12...HEAD [Unreleased]: https://github.com/an-tao/drogon/compare/v1.9.13...HEAD
[1.9.13]: https://github.com/an-tao/drogon/compare/v1.9.12...v1.9.13
[1.9.12]: https://github.com/an-tao/drogon/compare/v1.9.11...v1.9.12 [1.9.12]: https://github.com/an-tao/drogon/compare/v1.9.11...v1.9.12
+15
View File
@@ -21,6 +21,21 @@
# ############################################################################## # ##############################################################################
include(FindPackageHandleStandardArgs) include(FindPackageHandleStandardArgs)
# On Apple Silicon Macs, Homebrew installs to /opt/homebrew instead of
# /usr/local (Intel Macs). Detect the prefix dynamically so cmake finds
# dependencies regardless of Mac architecture.
if(APPLE)
execute_process(
COMMAND brew --prefix brotli
OUTPUT_VARIABLE HOMEBREW_BROTLI_PREFIX
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(HOMEBREW_BROTLI_PREFIX)
list(APPEND CMAKE_PREFIX_PATH ${HOMEBREW_BROTLI_PREFIX})
endif()
endif()
find_path(BROTLI_INCLUDE_DIR "brotli/decode.h") find_path(BROTLI_INCLUDE_DIR "brotli/decode.h")
find_library(BROTLICOMMON_LIBRARY NAMES brotlicommon brotlicommon-static) find_library(BROTLICOMMON_LIBRARY NAMES brotlicommon brotlicommon-static)
+15
View File
@@ -5,6 +5,21 @@
# HIREDIS_INCLUDE_DIRS - hiredis include directories # HIREDIS_INCLUDE_DIRS - hiredis include directories
# HIREDIS_LIBRARIES - libraries need to use hiredis # HIREDIS_LIBRARIES - libraries need to use hiredis
# On Apple Silicon Macs, Homebrew installs to /opt/homebrew instead of
# /usr/local (Intel Macs). Detect the prefix dynamically so cmake finds
# dependencies regardless of Mac architecture.
if(APPLE)
execute_process(
COMMAND brew --prefix hiredis
OUTPUT_VARIABLE HOMEBREW_HIREDIS_PREFIX
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(HOMEBREW_HIREDIS_PREFIX)
list(APPEND CMAKE_PREFIX_PATH ${HOMEBREW_HIREDIS_PREFIX})
endif()
endif()
if (HIREDIS_INCLUDE_DIRS AND HIREDIS_LIBRARIES) if (HIREDIS_INCLUDE_DIRS AND HIREDIS_LIBRARIES)
set(HIREDIS_FIND_QUIETLY TRUE) set(HIREDIS_FIND_QUIETLY TRUE)
set(Hiredis_FOUND TRUE) set(Hiredis_FOUND TRUE)
+16
View File
@@ -10,6 +10,22 @@
# false, do not try to use jsoncpp. # false, do not try to use jsoncpp.
# Jsoncpp_lib - The imported target library. # Jsoncpp_lib - The imported target library.
# On Apple Silicon Macs, Homebrew installs to /opt/homebrew instead of
# /usr/local (Intel Macs). Detect the prefix dynamically so cmake finds
# dependencies regardless of Mac architecture.
if(APPLE)
execute_process(
COMMAND brew --prefix jsoncpp
OUTPUT_VARIABLE HOMEBREW_JSONCPP_PREFIX
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(HOMEBREW_JSONCPP_PREFIX)
list(APPEND CMAKE_PREFIX_PATH ${HOMEBREW_JSONCPP_PREFIX})
endif()
endif()
# only look in default directories # only look in default directories
find_path(JSONCPP_INCLUDE_DIRS find_path(JSONCPP_INCLUDE_DIRS
NAMES json/json.h NAMES json/json.h
+15
View File
@@ -13,6 +13,21 @@
# SQLite3_FOUND - True if sqlite3 found. # SQLite3_FOUND - True if sqlite3 found.
# SQLite3_lib - The imported target library. # SQLite3_lib - The imported target library.
# On Apple Silicon Macs, Homebrew installs to /opt/homebrew instead of
# /usr/local (Intel Macs). Detect the prefix dynamically so cmake finds
# dependencies regardless of Mac architecture.
if(APPLE)
execute_process(
COMMAND brew --prefix sqlite3
OUTPUT_VARIABLE HOMEBREW_SQLITE3_PREFIX
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(HOMEBREW_SQLITE3_PREFIX)
list(APPEND CMAKE_PREFIX_PATH ${HOMEBREW_SQLITE3_PREFIX})
endif()
endif()
# Look for the header file. # Look for the header file.
find_path(SQLITE3_INCLUDE_DIRS NAMES sqlite3.h) find_path(SQLITE3_INCLUDE_DIRS NAMES sqlite3.h)
+15
View File
@@ -7,6 +7,21 @@
# PostgreSQL. # PostgreSQL.
# pg_lib - The imported target library. # pg_lib - The imported target library.
# On Apple Silicon Macs, Homebrew installs to /opt/homebrew instead of
# /usr/local (Intel Macs). Detect the prefix dynamically so cmake finds
# dependencies regardless of Mac architecture.
if(APPLE)
execute_process(
COMMAND brew --prefix libpq
OUTPUT_VARIABLE HOMEBREW_PG_PREFIX
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_QUIET
)
if(HOMEBREW_PG_PREFIX)
list(APPEND CMAKE_PREFIX_PATH ${HOMEBREW_PG_PREFIX})
endif()
endif()
find_package(PostgreSQL) find_package(PostgreSQL)
if(PostgreSQL_FOUND) if(PostgreSQL_FOUND)
set(PG_LIBRARIES ${PostgreSQL_LIBRARIES}) set(PG_LIBRARIES ${PostgreSQL_LIBRARIES})
+13 -1
View File
@@ -42,7 +42,8 @@ std::string create::detail()
"create a plugin named class_name\n\n" "create a plugin named class_name\n\n"
"drogon_ctl create project <project_name> //" "drogon_ctl create project <project_name> //"
"create a project named project_name\n\n" "create a project named project_name\n\n"
"drogon_ctl create model <model_path> [-o <output path>] " "drogon_ctl create model <model_path> [-o <output path>] [ "
"--clear-output]"
"[--table=<table_name>] [-f]//" "[--table=<table_name>] [-f]//"
"create model classes in model_path\n"; "create model classes in model_path\n";
} }
@@ -55,3 +56,14 @@ void create::handleCommand(std::vector<std::string> &parameters)
parameters[0] = createObjName; parameters[0] = createObjName;
exeCommand(parameters); exeCommand(parameters);
} }
// Prevent clang-cl/lld-link from discarding DrObject<T>::alloc_ on Windows.
// On COFF targets, clang places the CRT initializer for the template static
// member in the same COMDAT group as the variable itself. When no code in the
// translation unit takes the address of alloc_ (clang inlines className()),
// the entire COMDAT is eligible for elimination — and lld-link's /OPT:REF
// removes it, so DrClassMap is never populated.
// Explicit template instantiation forces a strong (non-COMDAT) definition,
// which the linker must keep. This is a no-op on MSVC, GCC, and ELF targets
// (where .init_array entries are GC roots and are never discarded).
template class drogon::DrObject<drogon_ctl::create>;
@@ -469,3 +469,6 @@ void create_controller::createARestfulController(const std::string &className,
std::cout << "File name: " << ctlName << ".h and " << ctlName << ".cc" std::cout << "File name: " << ctlName << ".h and " << ctlName << ".cc"
<< std::endl; << std::endl;
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::create_controller>;
+3
View File
@@ -116,3 +116,6 @@ void create_filter::handleCommand(std::vector<std::string> &parameters)
createFilterSourceFile(oSourceFile, className, fileName); createFilterSourceFile(oSourceFile, className, fileName);
} }
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::create_filter>;
+195 -3
View File
@@ -164,6 +164,59 @@ bool drogon_ctl::ConvertMethod::shouldConvert(const std::string &tableName,
} // endif } // endif
} }
/**
* @brief Try to add an auto-detected FK relationship to the list.
*
* Checks for duplicates against existing relationships, creates a
* Relationship object from the FK info, and appends it to the list.
* User-configured relationships always take priority.
*
* @param allRelationships The mutable vector of relationships.
* @param originalTable The table containing the FK column.
* @param fkColumn The FK column name.
* @param referencedTable The table referenced by the FK.
* @param referencedColumn The column referenced by the FK.
* @param normalizeNames If true, apply toLower() to table names.
*/
static void tryAddAutoRelationship(std::vector<Relationship> &allRelationships,
const std::string &originalTable,
const std::string &fkColumn,
const std::string &referencedTable,
const std::string &referencedColumn,
bool normalizeNames)
{
for (const auto &r : allRelationships)
{
if (r.originalKey() == fkColumn &&
r.targetTableName() == referencedTable)
{
return; // Already exists in user config
}
}
Json::Value relJson;
relJson["type"] = "has one";
relJson["original_table_name"] =
normalizeNames ? toLower(originalTable) : originalTable;
relJson["original_key"] = fkColumn;
relJson["target_table_name"] =
normalizeNames ? toLower(referencedTable) : referencedTable;
relJson["target_key"] = referencedColumn;
relJson["enable_reverse"] = true;
try
{
Relationship autoRel(relJson);
allRelationships.push_back(autoRel);
std::cout << " Auto-detected FK: " << originalTable << "."
<< fkColumn << " -> " << referencedTable << "."
<< referencedColumn << std::endl;
}
catch (const std::runtime_error &e)
{
std::cerr << "Warning: Could not create auto-relationship: " << e.what()
<< std::endl;
}
}
#if USE_POSTGRESQL #if USE_POSTGRESQL
void create_model::createModelClassFromPG( void create_model::createModelClassFromPG(
const std::string &path, const std::string &path,
@@ -182,8 +235,9 @@ void create_model::createModelClassFromPG(
data["primaryKeyName"] = ""; data["primaryKeyName"] = "";
data["dbName"] = dbname_; data["dbName"] = dbname_;
data["rdbms"] = std::string("postgresql"); data["rdbms"] = std::string("postgresql");
data["relationships"] = relationships;
data["convertMethods"] = convertMethods; data["convertMethods"] = convertMethods;
// Start with user-configured relationships (mutable copy)
std::vector<Relationship> allRelationships(relationships);
if (schema != "public") if (schema != "public")
{ {
data["schema"] = schema; data["schema"] = schema;
@@ -397,6 +451,42 @@ void create_model::createModelClassFromPG(
data["primaryKeyValNames"] = pkValNames; data["primaryKeyValNames"] = pkValNames;
} }
// Auto-detect foreign key relationships from database schema
*client << "SELECT "
"kcu.column_name AS fk_column, "
"ccu.table_name AS referenced_table, "
"ccu.column_name AS referenced_column "
"FROM information_schema.key_column_usage kcu "
"JOIN information_schema.referential_constraints rc "
"ON kcu.constraint_name = rc.constraint_name "
"AND kcu.constraint_schema = rc.constraint_schema "
"JOIN information_schema.constraint_column_usage ccu "
"ON rc.unique_constraint_name = ccu.constraint_name "
"AND rc.unique_constraint_schema = ccu.constraint_schema "
"WHERE kcu.table_name = $1 "
"AND kcu.table_schema = $2"
<< tableName << schema << Mode::Blocking >>
[&](bool isNull,
const std::string &fkColumn,
const std::string &referencedTable,
const std::string &referencedColumn) {
if (!isNull)
{
tryAddAutoRelationship(allRelationships,
tableName,
fkColumn,
referencedTable,
referencedColumn,
true);
}
} >>
[](const DrogonDbException &e) {
// FK detection is best-effort; don't fail if unsupported
std::cerr << "Note: FK auto-detection not available: "
<< e.base().what() << std::endl;
};
data["relationships"] = allRelationships;
data["columns"] = cols; data["columns"] = cols;
std::ofstream headerFile(path + "/" + className + ".h", std::ofstream::out); std::ofstream headerFile(path + "/" + className + ".h", std::ofstream::out);
std::ofstream sourceFile(path + "/" + className + ".cc", std::ofstream sourceFile(path + "/" + className + ".cc",
@@ -467,8 +557,9 @@ void create_model::createModelClassFromMysql(
data["primaryKeyName"] = ""; data["primaryKeyName"] = "";
data["dbName"] = dbname_; data["dbName"] = dbname_;
data["rdbms"] = std::string("mysql"); data["rdbms"] = std::string("mysql");
data["relationships"] = relationships;
data["convertMethods"] = convertMethods; data["convertMethods"] = convertMethods;
// Start with user-configured relationships (mutable copy)
std::vector<Relationship> allRelationships(relationships);
std::vector<ColumnInfo> cols; std::vector<ColumnInfo> cols;
int i = 0; int i = 0;
*client << "desc `" + tableName + "`" << Mode::Blocking >> *client << "desc `" + tableName + "`" << Mode::Blocking >>
@@ -593,6 +684,35 @@ void create_model::createModelClassFromMysql(
data["primaryKeyType"] = pkTypes; data["primaryKeyType"] = pkTypes;
data["primaryKeyValNames"] = pkValNames; data["primaryKeyValNames"] = pkValNames;
} }
// Auto-detect foreign key relationships from MySQL schema
*client << "SELECT COLUMN_NAME, REFERENCED_TABLE_NAME, "
"REFERENCED_COLUMN_NAME "
"FROM information_schema.KEY_COLUMN_USAGE "
"WHERE TABLE_SCHEMA = DATABASE() "
"AND TABLE_NAME = ? "
"AND REFERENCED_TABLE_NAME IS NOT NULL"
<< tableName << Mode::Blocking >>
[&](bool isNull,
const std::string &fkColumn,
const std::string &referencedTable,
const std::string &referencedColumn) {
if (!isNull)
{
tryAddAutoRelationship(allRelationships,
tableName,
fkColumn,
referencedTable,
referencedColumn,
true);
}
} >>
[](const DrogonDbException &e) {
std::cerr << "Note: FK auto-detection not available: "
<< e.base().what() << std::endl;
};
data["relationships"] = allRelationships;
data["columns"] = cols; data["columns"] = cols;
std::ofstream headerFile(path + "/" + className + ".h", std::ofstream::out); std::ofstream headerFile(path + "/" + className + ".h", std::ofstream::out);
std::ofstream sourceFile(path + "/" + className + ".cc", std::ofstream sourceFile(path + "/" + className + ".cc",
@@ -646,8 +766,9 @@ void create_model::createModelClassFromSqlite3(
data["primaryKeyName"] = ""; data["primaryKeyName"] = "";
data["dbName"] = std::string("sqlite3"); data["dbName"] = std::string("sqlite3");
data["rdbms"] = std::string("sqlite3"); data["rdbms"] = std::string("sqlite3");
data["relationships"] = relationships;
data["convertMethods"] = convertMethods; data["convertMethods"] = convertMethods;
// Start with user-configured relationships (mutable copy)
std::vector<Relationship> allRelationships(relationships);
std::vector<ColumnInfo> cols; std::vector<ColumnInfo> cols;
std::string sql = "PRAGMA table_info(" + tableName + ");"; std::string sql = "PRAGMA table_info(" + tableName + ");";
*client << sql << Mode::Blocking >> [&](const Result &result) { *client << sql << Mode::Blocking >> [&](const Result &result) {
@@ -774,6 +895,28 @@ void create_model::createModelClassFromSqlite3(
data["primaryKeyType"] = pkTypes; data["primaryKeyType"] = pkTypes;
data["primaryKeyValNames"] = pkValNames; data["primaryKeyValNames"] = pkValNames;
} }
// Auto-detect foreign key relationships from SQLite3 schema
std::string fkSql = "PRAGMA foreign_key_list(\"" + tableName + "\");";
*client << fkSql << Mode::Blocking >> [&](const Result &fkResult) {
for (auto &fkRow : fkResult)
{
auto referencedTable = fkRow["table"].as<std::string>();
auto fkColumn = fkRow["from"].as<std::string>();
auto referencedColumn = fkRow["to"].as<std::string>();
tryAddAutoRelationship(allRelationships,
tableName,
fkColumn,
referencedTable,
referencedColumn,
true);
}
} >> [](const DrogonDbException &e) {
std::cerr << "Note: FK auto-detection not available: "
<< e.base().what() << std::endl;
};
data["relationships"] = allRelationships;
data["columns"] = cols; data["columns"] = cols;
std::ofstream headerFile(path + "/" + className + ".h", std::ofstream::out); std::ofstream headerFile(path + "/" + className + ".h", std::ofstream::out);
std::ofstream sourceFile(path + "/" + className + ".cc", std::ofstream sourceFile(path + "/" + className + ".cc",
@@ -826,7 +969,42 @@ void create_model::createModel(const std::string &path,
auto restfulApiConfig = config["restful_api_controllers"]; auto restfulApiConfig = config["restful_api_controllers"];
auto relationships = getRelationships(config["relationships"]); auto relationships = getRelationships(config["relationships"]);
auto convertMethods = getConvertMethods(config["convert"]); auto convertMethods = getConvertMethods(config["convert"]);
drogon::utils::createPath(path); drogon::utils::createPath(path);
if (cleanupDirectory_)
{
std::cout << "Source files (*.h, *.cc) in '" << path
<< "' folder will be deleted, continue(y/n)?\n";
auto in = getchar();
(void)getchar(); // get the return key
if (in != 'Y' && in != 'y')
{
std::cout << "Abort!" << std::endl;
exit(0);
}
for (const auto &entry : std::filesystem::directory_iterator(path))
{
if (!entry.is_regular_file())
continue;
const std::filesystem::path &file = entry.path();
std::string ext = file.extension().string();
if (ext == ".h" || ext == ".cc")
{
std::cout << "Removing: " << file << "\n";
std::error_code ret;
std::filesystem::remove(file, ret);
if (ret)
{
std::cerr << "Failed to remove '" << file
<< "' : " << ret.message() << "\n";
}
}
}
}
if (dbType == "postgresql") if (dbType == "postgresql")
{ {
#if USE_POSTGRESQL #if USE_POSTGRESQL
@@ -1230,6 +1408,17 @@ void create_model::handleCommand(std::vector<std::string> &parameters)
++iter; ++iter;
} }
for (auto iter = parameters.begin(); iter != parameters.end(); ++iter)
{
if ((*iter) == "--clear-output")
{
cleanupDirectory_ = true;
forceOverwrite_ = true;
parameters.erase(iter);
break;
}
}
for (auto const &path : parameters) for (auto const &path : parameters)
{ {
createModel(path, singleModelName); createModel(path, singleModelName);
@@ -1414,3 +1603,6 @@ void create_model::createRestfulAPIController(
<< std::endl; << std::endl;
} }
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::create_model>;
+1
View File
@@ -430,5 +430,6 @@ class create_model : public DrObject<create_model>, public CommandHandler
std::string dbname_; std::string dbname_;
bool forceOverwrite_{false}; bool forceOverwrite_{false};
std::string outputPath_; std::string outputPath_;
bool cleanupDirectory_{false};
}; };
} // namespace drogon_ctl } // namespace drogon_ctl
+3
View File
@@ -116,3 +116,6 @@ void create_plugin::handleCommand(std::vector<std::string> &parameters)
createPluginSourceFile(oSourceFile, className, fileName); createPluginSourceFile(oSourceFile, className, fileName);
} }
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::create_plugin>;
+3
View File
@@ -141,3 +141,6 @@ void create_project::createProject(const std::string &projectName)
std::ofstream testCmakeFile("test/CMakeLists.txt", std::ofstream::out); std::ofstream testCmakeFile("test/CMakeLists.txt", std::ofstream::out);
newTestCmakeFile(testCmakeFile, projectName); newTestCmakeFile(testCmakeFile, projectName);
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::create_project>;
+3
View File
@@ -552,3 +552,6 @@ void create_view::newViewSourceFile(std::ofstream &file,
file << "return templ->genText(data);\n"; file << "return templ->genText(data);\n";
file << "}\n}\n"; file << "}\n}\n";
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::create_view>;
+3
View File
@@ -68,3 +68,6 @@ void help::handleCommand(std::vector<std::string> &parameters)
} }
} }
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::help>;
+3
View File
@@ -470,3 +470,6 @@ void press::outputResults()
<< std::endl; << std::endl;
exit(0); exit(0);
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::press>;
+3
View File
@@ -66,3 +66,6 @@ void version::handleCommand(std::vector<std::string> &parameters)
std::cout << " yaml-cpp: no\n"; std::cout << " yaml-cpp: no\n";
#endif #endif
} }
// See create.cc for rationale.
template class drogon::DrObject<drogon_ctl::version>;
+42
View File
@@ -148,4 +148,46 @@ class HttpCoroMiddleware : public DrObject<T>, public HttpMiddlewareBase
#endif #endif
/**
* @brief Simple middleware that tags OPTIONS requests
* @details It adds the attribute "drogon.customCORShandling" to the request, so
* that HttpServer does not handle CORS for them internally.
*
* This allows custom CORS handling via the path handlers.
* For example to restrict the origins, headers allowed, specify a max age to
* avoid OPTIONS on every request, etc.
*
* Just register it:
* 1. globally via
* app().registerMiddleware(std::make_shared<drogon::HttpOptionsMiddleware>())
* 2. on every path handlers that need non-default handling, with
* ADD_METHOD_TO(..., drogon::Options, "drogon::HttpOptionsMiddleware")
*/
template <class Derived, bool AutoCreation = true>
class HttpOptionsMiddlewareImpl
: public drogon::HttpMiddleware<Derived, AutoCreation>
{
public:
void invoke(const HttpRequestPtr &req,
MiddlewareNextCallback &&nextCb,
MiddlewareCallback &&mcb) override
{
// Tag OPTIONS
if (req->method() == drogon::HttpMethod::Options)
req->attributes()->insert("drogon.customCORShandling", true);
// continue with next middleware (no post-processing here)
nextCb(std::move(mcb));
}
};
class HttpOptionsMiddlewareAuto
: public HttpOptionsMiddlewareImpl<HttpOptionsMiddlewareAuto, true>
{
};
class HttpOptionsMiddleware
: public HttpOptionsMiddlewareImpl<HttpOptionsMiddleware, false>
{
};
} // namespace drogon } // namespace drogon
+50
View File
@@ -159,6 +159,9 @@ class DROGON_EXPORT HttpRequest
*/ */
virtual void removeHeader(std::string key) = 0; virtual void removeHeader(std::string key) = 0;
// Clear all HTTP headers
virtual void clearHeaders() = 0;
/// Get the cookie string identified by the field parameter /// Get the cookie string identified by the field parameter
virtual const std::string &getCookie(const std::string &field) const = 0; virtual const std::string &getCookie(const std::string &field) const = 0;
@@ -415,6 +418,8 @@ class DROGON_EXPORT HttpRequest
virtual void setMethod(const HttpMethod method) = 0; virtual void setMethod(const HttpMethod method) = 0;
/// Set the path of the request /// Set the path of the request
/// @note The path is automatically encoded. use
/// @c setPathEncode(false) to avoid this.
virtual void setPath(const std::string &path) = 0; virtual void setPath(const std::string &path) = 0;
virtual void setPath(std::string &&path) = 0; virtual void setPath(std::string &&path) = 0;
@@ -432,6 +437,20 @@ class DROGON_EXPORT HttpRequest
virtual void setParameter(const std::string &key, virtual void setParameter(const std::string &key,
const std::string &value) = 0; const std::string &value) = 0;
/**
* Set the parameter to the query,
* regardless of the HTTP method or content type
*/
virtual void setQueryParameter(const std::string &key,
const std::string &value) = 0;
/**
* Set the parameter to the request body.
* @warning The content type must be @c application/x-www-form-urlencoded
* or @c multipart/form-data
*/
virtual void setBodyParameter(const std::string &key,
const std::string &value) = 0;
/// Set or get the content type /// Set or get the content type
virtual void setContentTypeCode(const ContentType type) = 0; virtual void setContentTypeCode(const ContentType type) = 0;
@@ -501,6 +520,37 @@ class DROGON_EXPORT HttpRequest
return toRequest(std::forward<T>(obj)); return toRequest(std::forward<T>(obj));
} }
/*! \brief Check if the request is a CORS request.
* \details It should contain:
* - Origin: origination page
* \returns true if the Origin header is present
*/
inline bool isCorsRequest() const
{
// Check presence of required headers
return headers().find("origin") != headers().end();
}
/*! \brief Check if the request is a CORS pre-flight request.
* \details Check if the method of the request is OPTIONS and if it is
* a CORS pre-flight request.\n
* It should contain:
* - Origin: origination page
* - Access-Control-Request-Method: method to be used in the
* actual request
* \returns true if the method is OPTIONS and the required CORS pre-flight
* headers are present
*/
inline bool isCorsPreflightRequest() const
{
if (method() != HttpMethod::Options)
return false;
// Check presence of required headers
return isCorsRequest() &&
headers().find("access-control-request-method") !=
headers().end();
}
virtual bool isOnSecureConnection() const noexcept = 0; virtual bool isOnSecureConnection() const noexcept = 0;
virtual void setContentTypeString(const char *typeString, virtual void setContentTypeString(const char *typeString,
size_t typeStringLength) = 0; size_t typeStringLength) = 0;
+143 -2
View File
@@ -161,6 +161,12 @@ class DROGON_EXPORT HttpResponse
setCustomStatusCode(code, message.data(), message.length()); setCustomStatusCode(code, message.data(), message.length());
} }
/// Set whether the response should be compress.
virtual void setAllowCompression(bool allow) = 0;
/// Get whether the response allow compression.
virtual bool allowCompression() const = 0;
/// Get the creation timestamp of the response. /// Get the creation timestamp of the response.
virtual const trantor::Date &creationDate() const = 0; virtual const trantor::Date &creationDate() const = 0;
@@ -552,6 +558,141 @@ class DROGON_EXPORT HttpResponse
return toResponse(std::forward<T>(obj)); return toResponse(std::forward<T>(obj));
} }
/*! \brief Create an OPTIONS or CORS pre-flight response
* \details If the request is not an OPTIONS request, returns a NULL
* response\n
* If it is a generic OPTIONS request, returns a 204 No Content
* response with the Allow header\n
* If it is a CORS pre-flight request, returns a 204 No Content
* response with the CORS headers set
*
* Other status codes for CORS pre-flight answers:
* - 400 Bad Request: if the request is malformed (missing
* required headers)
* - 403 Forbidden: if the Origin is not allowed + reason
* in a X-Cors-Error header
* - 403 Forbidden: if one of the headers in
* Access-Control-Request-Headers is not allowed + reason in
* a X-Cors-Error header
* - 405 Method Not Allowed: if the requested method is
* not allowed
* \note CORS is a browser-side security mechanism.\n
* Do not rely on Origin for authentication/authorization:
* non-browser clients can spoof or omit it.\n
* Enforce access control independently.
* \param[in] request Drogon (OPTIONS) request
* \param[in] allowedHeaders Set of allowed headers (for
* Access-Control-Allow-Headers header)\n
* (headers allowed by the controller path
* handler)
* \param[in] originValidator Function to validate the Origin header value
* (allow the origin or not)\n
* If allowCredentials is true, originValidator
* _SHOULD_ enforce a strict allowlist
* \param[in] allowNullOrigin Should be true to accept the "Origin: null"
* header\n
* (set for local file:// pages, sandboxed
* iframes, opaque origins, data: URIs)
* \param[in] allowCredentials Should be true to add the header
* "Access-Control-Allow-Credentials: true"
* (controls whether the browser may include
* credentials such as cookies, HTTP auth, or
* client certificates)\n
* Note: Authorization (bearer) is not a
* credential header; allow it via
* allowedHeaders when needed
* \param[in] allowPNA Should be true to accept the header
* "Access-Control-Request-Private-Network"
* (when a page from a less private address
* space is trying to reach a more private
* one, like internet -> intranet)\n
* Note: specific to Chromium & derivatives
* (Edge, Opera, Brave, ...), not in Firefox
* or Safari
* \param[in] maxAgeSeconds If set, adds the "Access-Control-Max-Age"
* header with the given value (in seconds,
* how long the results of a preflight
* request can be cached by the navigator)
* \returns the OPTIONS or CORS pre-flight response, or a null pointer if
* the request is not an OPTIONS request
*/
static HttpResponsePtr newOptionsResponse(
const HttpRequestPtr &request,
const std::function<bool(std::string_view)> &originValidator = nullptr,
bool allowNullOrigin = false,
bool allowCredentials = false,
bool allowPNA = true,
std::optional<unsigned int> maxAgeSeconds = {},
const std::optional<std::set<std::string_view>> &allowedHeaders =
std::nullopt);
/*! \copydoc newOptionsResponse(const HttpRequestPtr&,
* const std::function<bool(std::string_view)>&,
* bool, bool, bool,
* std::optional<unsigned int>,
* const std::optional<std::set<std::string_view>>&)
* \remarks Helper when specifying the allowed headers, when other
* parameters may be default, to avoid having to specify them all
*/
inline static HttpResponsePtr newOptionsResponse(
const HttpRequestPtr &request,
const std::set<std::string_view> &allowedHeaders,
const std::function<bool(std::string_view)> &originValidator = nullptr,
bool allowNullOrigin = false,
bool allowCredentials = false,
bool allowPNA = true,
std::optional<unsigned int> maxAgeSeconds = {})
{
return newOptionsResponse(request,
originValidator,
allowNullOrigin,
allowCredentials,
allowPNA,
maxAgeSeconds,
allowedHeaders);
}
/*! \brief Add CORS headers to a response
* \details Adds the CORS headers to a response for a normal request (a
* CORS request but not a CORS preflight request):
* - does nothing if it's an OPTIONS request, or
* - if it's not a CORS request, or
* - if it's a CORS preflight request
* Else:
* - adds Access-Control-Allow-Origin (if not yet present)
* - adds Origin to the Vary header,
* - sets or clears Access-Control-Allow-Credentials (if
* allowCredentials is set)
* - completes Access-Control-Expose-Headers
* \param[in] request Drogon request (to get Origin)
* \param[in] allowCredentials If set and true, adds the
* "Access-Control-Allow-Credentials: true
* header"\n
* If set and false, removes the
* "Access-Control-Allow-Credentials" header\n
* If not set, leaves the
* "Access-Control-Allow-Credentials" header
* untouched\n
* *MUST MATCH THE newOptionsResponse()
* PRE-FLIGHT RESPONSE VALUE*
* \param[in] exposedHeaders Set of exposed headers (for
* Access-Control-Expose-Headers header)\n
* These are the headers allowed to be exposed
* to javascript by the remote browser\n
* Note: they are *APPENDED* to any already
* present in the response, they are not
* REPLACED.\n
* This allows to complete them in the
* controller path handler.\n
* If you want to REPLACE them, remove the
* header before calling this function.
* \note may be use both in the controller path handler and in a
* pre-sending advice
*/
void addCorsHeaders(const HttpRequestPtr &request,
const std::set<std::string_view> &exposedHeaders = {},
const std::optional<bool> &allowCredentials = {});
/** /**
* @brief If the response is a file response (i.e. created by * @brief If the response is a file response (i.e. created by
* newFileResponse) returns the path on the filesystem. Otherwise a * newFileResponse) returns the path on the filesystem. Otherwise a
@@ -560,9 +701,9 @@ class DROGON_EXPORT HttpResponse
virtual const std::string &sendfileName() const = 0; virtual const std::string &sendfileName() const = 0;
/** /**
* @brief Returns the range of the file response as a pair ot size_t * @brief Returns the range of the file response as a pair of size_t
* (offset, length). Length of 0 means the entire file is sent. Behavior of * (offset, length). Length of 0 means the entire file is sent. Behavior of
* this function is undefined if the response if not a file response * this function is undefined if the response is not a file response
*/ */
using SendfileRange = std::pair<size_t, size_t>; // { offset, length } using SendfileRange = std::pair<size_t, size_t>; // { offset, length }
virtual const SendfileRange &sendfileRange() const = 0; virtual const SendfileRange &sendfileRange() const = 0;
+12
View File
@@ -195,6 +195,10 @@ enum HttpMethod
Delete, Delete,
Options, Options,
Patch, Patch,
Propfind,
Mkcol,
Copy,
Move,
Invalid Invalid
}; };
@@ -280,6 +284,14 @@ inline std::string_view to_string_view(drogon::HttpMethod method)
return "OPTIONS"; return "OPTIONS";
case drogon::HttpMethod::Patch: case drogon::HttpMethod::Patch:
return "PATCH"; return "PATCH";
case drogon::HttpMethod::Propfind:
return "PROPFIND";
case drogon::HttpMethod::Mkcol:
return "MKCOL";
case drogon::HttpMethod::Copy:
return "COPY";
case drogon::HttpMethod::Move:
return "MOVE";
default: default:
return "INVALID"; return "INVALID";
} }
+33
View File
@@ -56,6 +56,27 @@ class UploadFile
} }
} }
/// Constructor
/**
* @param data Pointer to the data
* @param len Data length in bytes
* @param fileName The file name provided to the server.
* @param itemName The item name on the browser form.
* @param contentType The Mime content type for the part
*/
explicit UploadFile(const void *data,
const size_t len,
const std::string &fileName = "memory.bin",
const std::string &itemName = "file",
ContentType contentType = CT_APPLICATION_OCTET_STREAM)
: data_(data),
len_(len),
fileName_(fileName),
itemName_(itemName),
contentType_(contentType)
{
}
const std::string &path() const const std::string &path() const
{ {
return path_; return path_;
@@ -76,7 +97,19 @@ class UploadFile
return contentType_; return contentType_;
} }
const void *data() const
{
return data_;
}
size_t dataLength() const
{
return len_;
}
private: private:
const void *data_ = nullptr;
size_t len_ = 0;
std::string path_; std::string path_;
std::string fileName_; std::string fileName_;
std::string itemName_; std::string itemName_;
@@ -11,6 +11,7 @@
#include <drogon/plugins/Plugin.h> #include <drogon/plugins/Plugin.h>
#include <trantor/utils/AsyncFileLogger.h> #include <trantor/utils/AsyncFileLogger.h>
#include <vector> #include <vector>
#include <regex>
namespace drogon namespace drogon
{ {
+162 -2
View File
@@ -124,6 +124,165 @@ DROGON_EXPORT std::set<std::string> splitStringToSet(
const std::string &str, const std::string &str,
const std::string &separator); const std::string &separator);
/*! \brief Compare two string_views for equality, ignoring case.
* \warning This is locale dependent
* \param[in] str1 The first string_view.
* \param[in] str2 The second string_view.
* \return true if the string_views are equal, ignoring case; false otherwise.
*/
inline bool ci_equals(std::string_view str1, std::string_view str2)
{
if (str1.size() != str2.size())
return false;
return std::equal(str1.begin(),
str1.end(),
str2.begin(),
[](unsigned char a, unsigned char b) {
return std::tolower(a) == std::tolower(b);
});
}
/*! \details Trim leading and trailing spaces and tabs from a string_view,
* modifying it.
* \param[in,out] str The string_view to trim.
* \return The trimmed string_view.
*/
inline std::string_view &trim_inplace(std::string_view &str)
{
auto pos = str.find_first_not_of(" \t");
// defeat Windows macro "min"
str.remove_prefix((std::min)(pos, str.size()));
if (str.empty())
return str;
pos = str.find_last_not_of(" \t");
str.remove_suffix(str.size() - pos - 1);
return str;
}
/*! \brief Trim leading and trailing spaces and tabs from a string_view.
* \param[in] str The string_view to trim.
* \return A string_view with leading and trailing spaces and tabs removed.
*/
inline std::string_view trim(std::string_view str)
{
return trim_inplace(str);
}
/*! \brief Trim leading and trailing spaces and tabs from a rvalue string.
* \param[in] str The string to trim.
* \return The string with leading and trailing spaces and tabs removed.
*/
inline std::string trim(std::string &&str)
{
auto pos = str.find_last_not_of(" \t");
if (pos == std::string::npos)
return {};
str.resize(pos + 1);
pos = str.find_first_not_of(" \t");
if (pos > 0)
str.erase(0, pos);
return str;
}
/*! \brief Split a string_view into a vector of string_views.
* \param[in] str The string_view to split.
* \param[in] separator The separator to use for splitting.
* \param[in] trimValues Whether to trim whitespace from the resulting
* string_views.
* \param[in] acceptEmptyString Whether to include empty strings in the result.
* \return A vector of string_views obtained by splitting the input
* string_view.
*/
inline std::vector<std::string_view> splitStringView(
std::string_view str,
std::string_view separator,
bool trimValues = true,
bool acceptEmptyString = false)
{
std::vector<std::string_view> result;
if (separator.empty())
{
if (trimValues)
trim_inplace(str);
if (acceptEmptyString || !str.empty())
result.push_back(str);
return result;
}
size_t start = 0;
size_t end = 0;
while ((end = str.find(separator, start)) != std::string_view::npos)
{
auto token = str.substr(start, end - start);
if (trimValues)
trim_inplace(token);
if (acceptEmptyString || !token.empty())
result.push_back(token);
start = end + separator.size();
}
auto token = str.substr(start);
if (trimValues)
trim_inplace(token);
if (acceptEmptyString || !token.empty())
{
result.push_back(token);
}
return result;
}
/*! \brief Split a string_view into a set of string_views.
* \copyparams splitStringView
* \return A set of (unique) string_views obtained by splitting the input
* string_view.
* \note Uniqueness is case-sensitive: "A" and "a" are considered different
* values.
*/
inline std::set<std::string_view> splitStringViewToSet(
std::string_view str,
std::string_view separator,
bool trimValues = true,
bool acceptEmptyString = false)
{
auto v = splitStringView(str, separator, trimValues, acceptEmptyString);
return std::set<std::string_view>(v.begin(), v.end());
}
/*! \brief Join a vector of string_view into a string.
* \param[in] strs The vector of string_views to join.
* \param[in] separator The separator to use between string_views.
* \return A single string obtained by joining the input string_views with the
* specified separator.
* \note Empty values are skipped.
*/
inline std::string joinStringViews(const std::vector<std::string_view> &strs,
std::string_view separator)
{
std::string result;
for (std::string_view str : strs)
{
if (trim_inplace(str).empty())
continue;
if (!result.empty())
result.append(separator);
result.append(str);
}
return result;
}
/*! \brief Join a set of string_view into a string.
* \param[in] strs The set of string_views to join.
* \param[in] separator The separator to use between string_views.
* \return A single string obtained by joining the input string_views with the
* specified separator.
* \note Empty values are skipped.
*/
inline std::string joinStringViews(const std::set<std::string_view> &strs,
std::string_view separator)
{
return joinStringViews(std::vector<std::string_view>{strs.begin(),
strs.end()},
separator);
}
/// Get UUID string. /// Get UUID string.
DROGON_EXPORT std::string getUuid(bool lowercase = true); DROGON_EXPORT std::string getUuid(bool lowercase = true);
@@ -497,7 +656,8 @@ T fromString(const std::string &p) noexcept(false)
// ("1a" should not return 1) // ("1a" should not return 1)
if (pos != p.size()) if (pos != p.size())
throw std::invalid_argument("Invalid value"); throw std::invalid_argument("Invalid value");
if ((v < static_cast<long double>((std::numeric_limits<T>::min)())) || if ((v <
static_cast<long double>((std::numeric_limits<T>::lowest)())) ||
(v > static_cast<long double>((std::numeric_limits<T>::max)()))) (v > static_cast<long double>((std::numeric_limits<T>::max)())))
throw std::out_of_range("Value out of range"); throw std::out_of_range("Value out of range");
return static_cast<T>(v); return static_cast<T>(v);
@@ -516,7 +676,7 @@ T fromString(const std::string &p) noexcept(false)
// throw if the whole string could not be parsed // throw if the whole string could not be parsed
// ("1a" should not return 1) // ("1a" should not return 1)
if (!ss.eof()) if (!ss.eof())
std::runtime_error("Bad type conversion"); throw std::runtime_error("Bad type conversion");
} }
return value; return value;
} }
@@ -57,12 +57,6 @@ void HttpConnectionLimit::releaseConnection(
const trantor::TcpConnectionPtr &conn) const trantor::TcpConnectionPtr &conn)
{ {
assert(!conn->connected()); assert(!conn->connected());
if (!conn->hasContext())
{
// If the connection is connected to the SSL port and then
// disconnected before the SSL handshake.
return;
}
connectionNum_.fetch_sub(1, std::memory_order_relaxed); connectionNum_.fetch_sub(1, std::memory_order_relaxed);
if (maxConnectionNumPerIP_ > 0) if (maxConnectionNumPerIP_ > 0)
{ {
@@ -63,6 +63,11 @@ void HttpControllersRouter::init(
initMiddlewaresAndCorsMethods(iter.second); initMiddlewaresAndCorsMethods(iter.second);
} }
for (auto &router : wsCtrlVector_)
{
initMiddlewaresAndCorsMethods(router);
}
for (auto &router : ctrlVector_) for (auto &router : ctrlVector_)
{ {
router.regex_ = std::regex(router.pathParameterPattern_, router.regex_ = std::regex(router.pathParameterPattern_,
@@ -85,6 +90,7 @@ void HttpControllersRouter::reset()
ctrlMap_.clear(); ctrlMap_.clear();
ctrlVector_.clear(); ctrlVector_.clear();
wsCtrlMap_.clear(); wsCtrlMap_.clear();
wsCtrlVector_.clear();
} }
std::vector<HttpHandlerInfo> HttpControllersRouter::getHandlersInfo() const std::vector<HttpHandlerInfo> HttpControllersRouter::getHandlersInfo() const
+73 -12
View File
@@ -215,6 +215,18 @@ void HttpRequestImpl::appendToBuffer(trantor::MsgBuffer *output) const
case Patch: case Patch:
output->append("PATCH "); output->append("PATCH ");
break; break;
case Propfind:
output->append("PROPFIND ");
break;
case Mkcol:
output->append("MKCOL ");
break;
case Copy:
output->append("COPY ");
break;
case Move:
output->append("MOVE ");
break;
default: default:
return; return;
} }
@@ -236,7 +248,7 @@ void HttpRequestImpl::appendToBuffer(trantor::MsgBuffer *output) const
} }
std::string content; std::string content;
if (passThrough_ && !query_.empty()) if (!query_.empty())
{ {
output->append("?"); output->append("?");
output->append(query_); output->append(query_);
@@ -323,21 +335,31 @@ void HttpRequestImpl::appendToBuffer(trantor::MsgBuffer *output) const
content.append(type.data(), type.length()); content.append(type.data(), type.length());
} }
content.append("\r\n\r\n"); content.append("\r\n\r\n");
std::ifstream infile(utils::toNativePath(file.path()),
std::ifstream::binary); if (file.data() && file.dataLength() > 0)
if (!infile)
{ {
LOG_ERROR << file.path() << " not found"; content.append((const char *)file.data(),
file.dataLength());
} }
else else
{ {
std::streambuf *pbuf = infile.rdbuf(); std::ifstream infile(utils::toNativePath(file.path()),
std::streamsize filesize = pbuf->pubseekoff(0, infile.end); std::ifstream::binary);
pbuf->pubseekoff(0, infile.beg); // rewind if (!infile)
std::string str; {
str.resize(filesize); LOG_ERROR << file.path() << " not found";
pbuf->sgetn(&str[0], filesize); }
content.append(std::move(str)); else
{
std::streambuf *pbuf = infile.rdbuf();
std::streamsize filesize =
pbuf->pubseekoff(0, infile.end);
pbuf->pubseekoff(0, infile.beg); // rewind
std::string str;
str.resize(filesize);
pbuf->sgetn(&str[0], filesize);
content.append(std::move(str));
}
} }
content.append("\r\n"); content.append("\r\n");
} }
@@ -648,6 +670,18 @@ const char *HttpRequestImpl::methodString() const
case Patch: case Patch:
result = "PATCH"; result = "PATCH";
break; break;
case Propfind:
result = "PROPFIND";
break;
case Mkcol:
result = "MKCOL";
break;
case Copy:
result = "COPY";
break;
case Move:
result = "MOVE";
break;
default: default:
break; break;
} }
@@ -683,6 +717,14 @@ bool HttpRequestImpl::setMethod(const char *start, const char *end)
{ {
method_ = Head; method_ = Head;
} }
else if (m == "COPY")
{
method_ = Copy;
}
else if (m == "MOVE")
{
method_ = Move;
}
else else
{ {
method_ = Invalid; method_ = Invalid;
@@ -693,6 +735,10 @@ bool HttpRequestImpl::setMethod(const char *start, const char *end)
{ {
method_ = Patch; method_ = Patch;
} }
else if (m == "MKCOL")
{
method_ = Mkcol;
}
else else
{ {
method_ = Invalid; method_ = Invalid;
@@ -718,6 +764,16 @@ bool HttpRequestImpl::setMethod(const char *start, const char *end)
method_ = Invalid; method_ = Invalid;
} }
break; break;
case 8:
if (m == "PROPFIND")
{
method_ = Propfind;
}
else
{
method_ = Invalid;
}
break;
default: default:
method_ = Invalid; method_ = Invalid;
break; break;
@@ -753,6 +809,11 @@ void HttpRequestImpl::reserveBodySize(size_t length)
{ {
// Store data of body to a temporary file // Store data of body to a temporary file
createTmpFile(); createTmpFile();
if (!content_.empty())
{
cacheFilePtr_->append(content_);
content_.clear();
}
} }
} }
+26
View File
@@ -351,6 +351,11 @@ class HttpRequestImpl : public HttpRequest
headers_.erase(lowerKey); headers_.erase(lowerKey);
} }
void clearHeaders() override
{
headers_.clear();
}
const std::string &getHeader(std::string field) const override const std::string &getHeader(std::string field) const override
{ {
std::transform(field.begin(), std::transform(field.begin(),
@@ -408,6 +413,27 @@ class HttpRequestImpl : public HttpRequest
parameters_[key] = value; parameters_[key] = value;
} }
void setQueryParameter(const std::string &key,
const std::string &value) override
{
if (!query_.empty())
{
query_.append("&");
}
query_.append(utils::urlEncodeComponent(key));
query_.append("=");
query_.append(utils::urlEncodeComponent(value));
}
void setBodyParameter(const std::string &key,
const std::string &value) override
{
assert(contentType_ == CT_MULTIPART_FORM_DATA ||
contentType_ == CT_APPLICATION_X_FORM);
flagForParsingParameters_ = true;
parameters_[key] = value;
}
const std::string &getContent() const const std::string &getContent() const
{ {
return content_; return content_;
+220
View File
@@ -54,6 +54,16 @@ static inline HttpResponsePtr genHttpResponse(const std::string &viewName,
} }
} // namespace drogon } // namespace drogon
void HttpResponseImpl::setAllowCompression(bool allow)
{
allowCompression_ = allow;
}
bool HttpResponseImpl::allowCompression() const
{
return allowCompression_;
}
HttpResponsePtr HttpResponse::newHttpResponse() HttpResponsePtr HttpResponse::newHttpResponse()
{ {
auto res = std::make_shared<HttpResponseImpl>(k200OK, CT_TEXT_HTML); auto res = std::make_shared<HttpResponseImpl>(k200OK, CT_TEXT_HTML);
@@ -473,6 +483,210 @@ HttpResponsePtr HttpResponse::newAsyncStreamResponse(
return resp; return resp;
} }
HttpResponsePtr HttpResponse::newOptionsResponse(
const HttpRequestPtr &request,
const std::function<bool(std::string_view)> &originValidator,
bool allowNullOrigin,
bool allowCredentials,
bool allowPNA,
std::optional<unsigned int> maxAgeSeconds,
const std::optional<std::set<std::string_view>> &allowedHeaders)
{
if (!request || (request->method() != HttpMethod::Options))
return {};
// Allowed methods, set by drogon::HttpOptionsMiddlewareImpl
auto methods =
request->attributes()->get<std::string>("drogon.corsMethods");
if (methods.empty())
methods = "OPTIONS";
auto response = newHttpResponse(HttpStatusCode::k204NoContent,
drogon::ContentType::CT_NONE);
// Disable HTTP caching for OPTIONS responses
response->addHeader("Cache-Control"s, "no-store"s);
// Vary on Origin for bad proxies that do not respect no-store or want
// Pragma: no-cache instead
response->addHeader("Vary"s, "Origin");
// Generic OPTIONS response
if (!request->isCorsPreflightRequest())
{
response->addHeader("Allow", methods);
return response;
}
// CORS pre-flight response
std::string_view origin = drogon::utils::trim(request->getHeader("Origin"));
if (origin.empty())
{
response->setStatusCode(HttpStatusCode::k400BadRequest);
response->addHeader("X-Cors-Error",
"invalid empty Origin"); // diagnose help
return response;
}
// Check whether null origin is allowed (file://, sandboxed iframes, etc.)
if (drogon::utils::ci_equals(origin, "null") && !allowNullOrigin)
{
response->setStatusCode(HttpStatusCode::k403Forbidden);
response->addHeader("X-Cors-Error",
"null Origin not allowed"); // diagnose help
return response;
}
// Check whether the origin is allowed
if (originValidator && !originValidator(origin))
{
response->setStatusCode(HttpStatusCode::k403Forbidden);
response->addHeader("X-Cors-Error",
"origin not allowed"); // diagnose help
return response;
}
// Reflect the origin (acts like '*', that is forbidden when
// allowCredentials is true)
response->addHeader("Access-Control-Allow-Origin", std::string(origin));
response->addHeader("Access-Control-Allow-Methods", methods);
// Check requested method
// Policy: explicitly fail preflight with 40x + diagnostic header rather
// than silently returning allowed methods
auto acrMethod = drogon::utils::trim(
request->getHeader("Access-Control-Request-Method"));
if (acrMethod.empty())
{
response->setStatusCode(HttpStatusCode::k400BadRequest);
response->addHeader(
"X-Cors-Error",
"invalid empty Access-Control-Request-Method"); // diagnose help
return response;
}
const auto allowedMethods = drogon::utils::splitStringView(methods, ",");
if (std::find_if(allowedMethods.begin(),
allowedMethods.end(),
[&acrMethod](const std::string_view &method) {
return drogon::utils::ci_equals(method, acrMethod);
}) == allowedMethods.end())
{
response->setStatusCode(HttpStatusCode::k405MethodNotAllowed);
response->addHeader("Allow",
methods); // failing CORS pre-flight with 405 must
// also return the Allow header
response->addHeader("X-Cors-Error",
"method not allowed: "s.append(
acrMethod)); // diagnose help
return response;
}
// Allowed headers (intersection with requested ones on success, all allowed
// on error) Note: Browsers typically include only non-safelisted headers in
// Access-Control-Request-Headers We validate strictly against
// allowedHeaders Policy: explicitly fail preflight with 403 + diagnostic
// header rather than silently omitting forbidden CORS headers
auto requestedHeaders = drogon::utils::splitStringViewToSet(
request->getHeader("Access-Control-Request-Headers"), ",");
if (allowedHeaders.has_value())
{
auto &validHeaders = allowedHeaders.value();
if (requestedHeaders.empty()) // noisy, but helpful for diagnosis
requestedHeaders = {validHeaders.begin(), validHeaders.end()};
else
{
for (auto it = requestedHeaders.begin();
it != requestedHeaders.end();)
{
auto &reqHeader = *it;
if (std::find_if(validHeaders.begin(),
validHeaders.end(),
[&reqHeader](const std::string_view &header) {
return drogon::utils::ci_equals(
reqHeader,
drogon::utils::trim(header));
}) != validHeaders.end())
{
++it;
continue;
}
response->setStatusCode(
HttpStatusCode::k403Forbidden); // Forbidden header
response->addHeader("X-Cors-Error",
"disallowed header: "s.append(
reqHeader)); // diagnose help
// report all allowed headers to help diagnosing what's
// wrong
requestedHeaders = {validHeaders.begin(), validHeaders.end()};
break;
}
}
}
if (!requestedHeaders.empty())
response->addHeader("Access-Control-Allow-Headers",
drogon::utils::joinStringViews(requestedHeaders,
","));
if (response->statusCode() == HttpStatusCode::k403Forbidden)
return response;
// Allow credentials
if (allowCredentials)
response->addHeader("Access-Control-Allow-Credentials", "true");
// Chromium-based browsers require this header to allow Private Network
// Access requests
if (allowPNA &&
drogon::utils::ci_equals(request->getHeader(
"Access-Control-Request-Private-Network"),
"true"))
response->addHeader("Access-Control-Allow-Private-Network", "true");
// Set a max age only on success
if (maxAgeSeconds.has_value())
response->addHeader("Access-Control-Max-Age",
std::to_string(maxAgeSeconds.value()));
return response;
}
void HttpResponse::addCorsHeaders(
const HttpRequestPtr &request,
const std::set<std::string_view> &exposedHeaders,
const std::optional<bool> &allowCredentials)
{
if (!request || !request->isCorsRequest() ||
request->isCorsPreflightRequest())
return;
// add/set Origin to the Vary header (needed for cache proxies)
auto vary = drogon::utils::splitStringViewToSet(getHeader("Vary"), ",");
if (std::find_if(vary.begin(), vary.end(), [](const auto &val) {
return drogon::utils::ci_equals(val, "Origin");
}) == vary.end())
{
vary.insert("Origin");
addHeader("Vary", drogon::utils::joinStringViews(vary, ","));
}
// add _MISSING_ CORS header - do not overwrite existing one
if (headers().find("access-control-allow-origin") == headers().end())
addHeader("Access-Control-Allow-Origin",
std::string(
drogon::utils::trim(request->getHeader("Origin"))));
// set (or append) exposed headers
if (!exposedHeaders.empty())
{
auto exposed = drogon::utils::splitStringViewToSet(
getHeader("Access-Control-Expose-Headers"), ",");
bool changed = false;
for (auto &header : exposedHeaders)
{
if (std::find_if(exposed.begin(),
exposed.end(),
[&header](const auto &val) {
return drogon::utils::ci_equals(val, header);
}) != exposed.end())
continue;
exposed.insert(header);
changed = true;
}
if (changed)
addHeader("Access-Control-Expose-Headers",
drogon::utils::joinStringViews(exposed, ","));
}
if (!allowCredentials.has_value())
return;
if (allowCredentials.value())
addHeader("Access-Control-Allow-Credentials", "true");
else
removeHeader("Access-Control-Allow-Credentials");
}
void HttpResponseImpl::makeHeaderString(trantor::MsgBuffer &buffer) void HttpResponseImpl::makeHeaderString(trantor::MsgBuffer &buffer)
{ {
buffer.ensureWritableBytes(128); buffer.ensureWritableBytes(128);
@@ -960,6 +1174,12 @@ void HttpResponseImpl::parseJson() const
bool HttpResponseImpl::shouldBeCompressed() const bool HttpResponseImpl::shouldBeCompressed() const
{ {
// If the developer said "No" stop immediately.
if (!allowCompression_)
{
return false;
}
if (streamCallback_ || asyncStreamCallback_ || !sendfileName_.empty() || if (streamCallback_ || asyncStreamCallback_ || !sendfileName_.empty() ||
contentType() >= CT_APPLICATION_OCTET_STREAM || contentType() >= CT_APPLICATION_OCTET_STREAM ||
getBody().length() < 1024 || getBody().length() < 1024 ||
+6
View File
@@ -463,6 +463,12 @@ class DROGON_EXPORT HttpResponseImpl : public HttpResponse
} }
private: private:
bool allowCompression_{true};
void setAllowCompression(bool allow) override;
bool allowCompression() const override;
void setBody(const char *body, size_t len) override void setBody(const char *body, size_t len) override
{ {
bodyPtr_ = std::make_shared<HttpMessageStringViewBody>(body, len); bodyPtr_ = std::make_shared<HttpMessageStringViewBody>(body, len);
+22 -6
View File
@@ -17,6 +17,8 @@
#include <trantor/utils/Logger.h> #include <trantor/utils/Logger.h>
#include <trantor/utils/MsgBuffer.h> #include <trantor/utils/MsgBuffer.h>
#include <algorithm> #include <algorithm>
#include <cerrno>
#include <cstdlib>
using namespace trantor; using namespace trantor;
using namespace drogon; using namespace drogon;
@@ -129,7 +131,16 @@ bool HttpResponseParser::parseResponse(MsgBuffer *buf)
// LOG_INFO << "content len=" << len; // LOG_INFO << "content len=" << len;
if (!len.empty()) if (!len.empty())
{ {
leftBodyLength_ = static_cast<size_t>(std::stoull(len)); try
{
leftBodyLength_ =
static_cast<size_t>(std::stoull(len));
}
catch (...)
{
// Malformed Content-Length from peer.
return false;
}
status_ = HttpResponseParseStatus::kExpectBody; status_ = HttpResponseParseStatus::kExpectBody;
} }
else else
@@ -242,12 +253,17 @@ bool HttpResponseParser::parseResponse(MsgBuffer *buf)
const char *crlf = buf->findCRLF(); const char *crlf = buf->findCRLF();
if (crlf) if (crlf)
{ {
// chunk length line
std::string len(buf->peek(), crlf - buf->peek()); std::string len(buf->peek(), crlf - buf->peek());
char *end; errno = 0;
currentChunkLength_ = strtol(len.c_str(), &end, 16); char *end = nullptr;
// LOG_TRACE << "chun length : " << unsigned long long parsed =
// currentChunkLength_; std::strtoull(len.c_str(), &end, 16);
if (errno == ERANGE || end == len.c_str() ||
(*end != '\0' && *end != ';'))
{
return false;
}
currentChunkLength_ = static_cast<size_t>(parsed);
if (currentChunkLength_ != 0) if (currentChunkLength_ != 0)
{ {
status_ = HttpResponseParseStatus::kExpectChunkBody; status_ = HttpResponseParseStatus::kExpectChunkBody;
+19 -12
View File
@@ -130,15 +130,20 @@ void HttpServer::onConnection(const TcpConnectionPtr &conn)
else if (conn->disconnected()) else if (conn->disconnected())
{ {
LOG_TRACE << "conn disconnected!"; LOG_TRACE << "conn disconnected!";
HttpConnectionLimit::instance().releaseConnection(conn);
auto requestParser = conn->getContext<HttpRequestParser>(); auto requestParser = conn->getContext<HttpRequestParser>();
if (requestParser) if (requestParser)
{ {
// NOTE: if tls handshake fails, `onConnection()` will only be
// called once with a broken conn. So we only call
// `releaseConnection()` for conn with context.
// Never call `conn->clearContext()` in other places
HttpConnectionLimit::instance().releaseConnection(conn);
if (requestParser->webSocketConn()) if (requestParser->webSocketConn())
{ {
requestParser->webSocketConn()->onClose(); requestParser->webSocketConn()->onClose();
} }
else if (requestParser->requestImpl()->isStreamMode()) else if (requestParser->requestImpl()->streamStatus() ==
ReqStreamStatus::Open)
{ {
requestParser->requestImpl()->streamError( requestParser->requestImpl()->streamError(
std::make_exception_ptr( std::make_exception_ptr(
@@ -206,13 +211,9 @@ void HttpServer::onMessage(const TcpConnectionPtr &conn, MsgBuffer *buf)
statusCodeToString(code).data())); statusCodeToString(code).data()));
} }
buf->retrieveAll(); buf->retrieveAll();
// NOTE: should we call conn->forceClose() instead? // stop parser to ignore following illegal data from client
// Calling shutdown() handles socket more elegantly. requestParser->stop();
conn->shutdown(); conn->shutdown();
// We have to call clearContext() here in order to ignore following
// illegal data from client
conn->clearContext();
requestParser->reset();
return; return;
} }
if (parseRes == 0) if (parseRes == 0)
@@ -576,12 +577,18 @@ void HttpServer::requestPassMiddlewares(const HttpRequestImplPtr &req,
template <typename Pack> template <typename Pack>
void HttpServer::requestPreHandling(const HttpRequestImplPtr &req, Pack &&pack) void HttpServer::requestPreHandling(const HttpRequestImplPtr &req, Pack &&pack)
{ {
// Handle CORS preflight request, except when custom handling is desired
if (req->method() == Options) if (req->method() == Options)
{ {
handleHttpOptions(req, if (!req->attributes()->get<bool>("drogon.customCORShandling"))
*pack.binderPtr->corsMethods_, {
std::move(pack.callback)); handleHttpOptions(req,
return; *pack.binderPtr->corsMethods_,
std::move(pack.callback));
return;
}
req->attributes()->insert("drogon.corsMethods",
*pack.binderPtr->corsMethods_);
} }
// pre-handling aop // pre-handling aop
+2 -25
View File
@@ -18,6 +18,7 @@
#include "HttpFileImpl.h" #include "HttpFileImpl.h"
#include <drogon/MultiPart.h> #include <drogon/MultiPart.h>
#include <drogon/utils/Utilities.h> #include <drogon/utils/Utilities.h>
#include "utils/ParsingUtils.h"
#include <drogon/config.h> #include <drogon/config.h>
#include <algorithm> #include <algorithm>
#include <fcntl.h> #include <fcntl.h>
@@ -29,6 +30,7 @@
#endif #endif
using namespace drogon; using namespace drogon;
using drogon::utils::parseLine;
const std::vector<HttpFile> &MultiPartParser::getFiles() const const std::vector<HttpFile> &MultiPartParser::getFiles() const
{ {
@@ -87,31 +89,6 @@ int MultiPartParser::parse(const HttpRequestPtr &req)
return parse(req, contentType.data() + (pos + 9), pos2 - (pos + 9)); return parse(req, contentType.data() + (pos + 9), pos2 - (pos + 9));
} }
static std::pair<std::string_view, std::string_view> parseLine(
const char *begin,
const char *end)
{
auto p = begin;
while (p != end)
{
if (*p == ':')
{
if (p + 1 != end && *(p + 1) == ' ')
{
return std::make_pair(std::string_view(begin, p - begin),
std::string_view(p + 2, end - p - 2));
}
else
{
return std::make_pair(std::string_view(begin, p - begin),
std::string_view(p + 1, end - p - 1));
}
}
++p;
}
return std::make_pair(std::string_view(), std::string_view());
}
int MultiPartParser::parseEntity(const HttpRequestPtr &req, int MultiPartParser::parseEntity(const HttpRequestPtr &req,
const char *begin, const char *begin,
const char *end) const char *end)
+4 -59
View File
@@ -14,41 +14,12 @@
#include "MultipartStreamParser.h" #include "MultipartStreamParser.h"
#include <cassert> #include <cassert>
#include "utils/ParsingUtils.h"
using namespace drogon; using namespace drogon;
using drogon::utils::parseLine;
static bool startsWith(const std::string_view &a, const std::string_view &b) using drogon::utils::startsWith;
{ using drogon::utils::startsWithIgnoreCase;
if (a.size() < b.size())
{
return false;
}
for (size_t i = 0; i < b.size(); i++)
{
if (a[i] != b[i])
{
return false;
}
}
return true;
}
static bool startsWithIgnoreCase(const std::string_view &a,
const std::string_view &b)
{
if (a.size() < b.size())
{
return false;
}
for (size_t i = 0; i < b.size(); i++)
{
if (::tolower(a[i]) != ::tolower(b[i]))
{
return false;
}
}
return true;
}
MultipartStreamParser::MultipartStreamParser(const std::string &contentType) MultipartStreamParser::MultipartStreamParser(const std::string &contentType)
{ {
@@ -86,32 +57,6 @@ MultipartStreamParser::MultipartStreamParser(const std::string &contentType)
crlfDashBoundary_ = crlf_ + dash_ + boundary_; crlfDashBoundary_ = crlf_ + dash_ + boundary_;
} }
// TODO: same function in HttpRequestParser.cc
static std::pair<std::string_view, std::string_view> parseLine(
const char *begin,
const char *end)
{
auto p = begin;
while (p != end)
{
if (*p == ':')
{
if (p + 1 != end && *(p + 1) == ' ')
{
return std::make_pair(std::string_view(begin, p - begin),
std::string_view(p + 2, end - p - 2));
}
else
{
return std::make_pair(std::string_view(begin, p - begin),
std::string_view(p + 1, end - p - 1));
}
}
++p;
}
return std::make_pair(std::string_view(), std::string_view());
}
void drogon::MultipartStreamParser::parse( void drogon::MultipartStreamParser::parse(
const char *data, const char *data,
size_t length, size_t length,
+94 -29
View File
@@ -17,10 +17,49 @@
#include <dirent.h> #include <dirent.h>
#include <dlfcn.h> #include <dlfcn.h>
#include <fstream> #include <fstream>
#include <sstream>
#include <sys/types.h> #include <sys/types.h>
#include <sys/wait.h>
#include <trantor/utils/Logger.h> #include <trantor/utils/Logger.h>
#include <unistd.h> #include <unistd.h>
// Safe exec helper: runs a program with explicit argv, no shell involved.
// Returns the exit status, or -1 on fork/exec failure.
static int safeExec(const std::vector<std::string> &args)
{
if (args.empty())
return -1;
std::vector<char *> argv;
argv.reserve(args.size() + 1);
for (auto &a : args)
argv.push_back(const_cast<char *>(a.c_str()));
argv.push_back(nullptr);
pid_t pid = fork();
if (pid == -1)
{
perror("fork");
return -1;
}
if (pid == 0)
{
// Child: replace image with the target program.
execvp(argv[0], argv.data());
// execvp only returns on error.
perror("execvp");
_exit(127);
}
// Parent: wait for child.
int status = 0;
if (waitpid(pid, &status, 0) == -1)
{
perror("waitpid");
return -1;
}
return WIFEXITED(status) ? WEXITSTATUS(status) : -1;
}
static void forEachFileIn( static void forEachFileIn(
const std::string &path, const std::string &path,
const std::function<void(const std::string &, const struct stat &)> &cb) const std::function<void(const std::string &, const struct stat &)> &cb)
@@ -153,22 +192,27 @@ void SharedLibManager::managerLibs()
else else
{ {
// generate source code and compile it. // generate source code and compile it.
std::string cmd = "drogon_ctl create view "; const std::string &outDir =
if (!outputPath_.empty()) !outputPath_.empty() ? outputPath_ : libPath;
{ std::vector<std::string> genArgs = {"drogon_ctl",
cmd.append(filename).append(" -o ").append( "create",
outputPath_); "view",
} filename,
else "-o",
{ outDir};
cmd.append(filename).append(" -o ").append(
libPath);
}
srcFile.append(".cc"); srcFile.append(".cc");
LOG_TRACE << cmd; LOG_TRACE << "drogon_ctl create view " << filename
auto r = system(cmd.c_str()); << " -o " << outDir;
// TODO: handle r auto r = safeExec(genArgs);
(void)(r); if (r != 0)
{
LOG_ERROR
<< "Failed to generate source code for "
<< filename;
dlStat.handle = oldHandle;
return;
}
dlStat.handle = dlStat.handle =
compileAndLoadLib(srcFile, oldHandle); compileAndLoadLib(srcFile, oldHandle);
} }
@@ -203,24 +247,45 @@ void *SharedLibManager::compileAndLoadLib(const std::string &sourceFile,
void *oldHld) void *oldHld)
{ {
LOG_TRACE << "src:" << sourceFile; LOG_TRACE << "src:" << sourceFile;
std::string cmd = COMPILER_COMMAND;
cmd.append(" ")
.append(sourceFile)
.append(" ")
.append(COMPILATION_FLAGS)
.append(" ")
.append(INCLUDING_DIRS);
if (std::string(COMPILER_ID).find("Clang") != std::string::npos)
cmd.append(" -shared -fPIC -undefined dynamic_lookup -o ");
else
cmd.append(" -shared -fPIC --no-gnu-unique -o ");
auto pos = sourceFile.rfind('.'); auto pos = sourceFile.rfind('.');
auto soFile = sourceFile.substr(0, pos); auto soFile = sourceFile.substr(0, pos);
soFile.append(".so"); soFile.append(".so");
cmd.append(soFile);
LOG_TRACE << cmd;
if (system(cmd.c_str()) == 0) // Build argv without invoking a shell so that metacharacters in
// sourceFile or soFile cannot be interpreted by /bin/sh.
std::vector<std::string> compileArgs;
compileArgs.push_back(COMPILER_COMMAND);
// COMPILATION_FLAGS and INCLUDING_DIRS are baked in at build time from
// trusted CMake variables; split them on whitespace into separate tokens.
auto splitIntoArgs = [&](const std::string &s) {
std::istringstream iss(s);
std::string token;
while (iss >> token)
compileArgs.push_back(token);
};
compileArgs.push_back(sourceFile);
splitIntoArgs(COMPILATION_FLAGS);
splitIntoArgs(INCLUDING_DIRS);
if (std::string(COMPILER_ID).find("Clang") != std::string::npos)
{
compileArgs.push_back("-shared");
compileArgs.push_back("-fPIC");
compileArgs.push_back("-undefined");
compileArgs.push_back("dynamic_lookup");
}
else
{
compileArgs.push_back("-shared");
compileArgs.push_back("-fPIC");
compileArgs.push_back("--no-gnu-unique");
}
compileArgs.push_back("-o");
compileArgs.push_back(soFile);
LOG_TRACE << COMPILER_COMMAND << " " << sourceFile << " ... -o " << soFile;
if (safeExec(compileArgs) == 0)
{ {
LOG_TRACE << "Compiled successfully:" << soFile; LOG_TRACE << "Compiled successfully:" << soFile;
return loadLib(soFile, oldHld); return loadLib(soFile, oldHld);
+52 -11
View File
@@ -155,9 +155,24 @@ bool isInteger(std::string_view str)
bool isBase64(std::string_view str) bool isBase64(std::string_view str)
{ {
for (auto c : str) if (str.empty())
if (!isBase64(c)) return false;
size_t padding = 0;
if (str.back() == '=')
padding++;
if (str.size() > 1 && str[str.size() - 2] == '=')
padding++;
for (size_t i = 0; i < str.size() - padding; ++i)
{
if (!isBase64(str[i]))
return false; return false;
}
if (padding > 0 && (str.size() % 4 != 0))
return false;
return true; return true;
} }
@@ -1018,6 +1033,35 @@ std::string gzipDecompress(const char *data, const size_t ndata)
} }
} }
static int formatHttpDate(char *buf, size_t len, const trantor::Date &date)
{
static const char *const weekdays[] = {
"Sun", "Mon", "Tue", "Wed", "Thu", "Fri", "Sat"};
static const char *const months[] = {"Jan",
"Feb",
"Mar",
"Apr",
"May",
"Jun",
"Jul",
"Aug",
"Sep",
"Oct",
"Nov",
"Dec"};
struct tm tm = date.tmStruct();
return snprintf(buf,
len,
"%s, %02d %s %04d %02d:%02d:%02d GMT",
weekdays[tm.tm_wday],
tm.tm_mday,
months[tm.tm_mon],
tm.tm_year + 1900,
tm.tm_hour,
tm.tm_min,
tm.tm_sec);
}
char *getHttpFullDate(const trantor::Date &date) char *getHttpFullDate(const trantor::Date &date)
{ {
static thread_local int64_t lastSecond = 0; static thread_local int64_t lastSecond = 0;
@@ -1029,9 +1073,7 @@ char *getHttpFullDate(const trantor::Date &date)
return lastTimeString; return lastTimeString;
} }
lastSecond = nowSecond; lastSecond = nowSecond;
date.toCustomFormattedString("%a, %d %b %Y %H:%M:%S GMT", formatHttpDate(lastTimeString, sizeof(lastTimeString), date);
lastTimeString,
sizeof(lastTimeString));
return lastTimeString; return lastTimeString;
} }
@@ -1039,8 +1081,6 @@ void dateToCustomFormattedString(const std::string &fmtStr,
std::string &str, std::string &str,
const trantor::Date &date) const trantor::Date &date)
{ {
auto nowSecond =
date.microSecondsSinceEpoch() / trantor::Date::MICRO_SECONDS_PER_SEC;
struct tm tm_LValue = date.tmStruct(); struct tm tm_LValue = date.tmStruct();
std::stringstream Out; std::stringstream Out;
Out.imbue(std::locale{"C"}); Out.imbue(std::locale{"C"});
@@ -1051,7 +1091,7 @@ void dateToCustomFormattedString(const std::string &fmtStr,
const std::string &getHttpFullDateStr(const trantor::Date &date) const std::string &getHttpFullDateStr(const trantor::Date &date)
{ {
static thread_local int64_t lastSecond = 0; static thread_local int64_t lastSecond = 0;
static thread_local std::string lastTimeString(128, 0); static thread_local std::string lastTimeString;
auto nowSecond = auto nowSecond =
date.microSecondsSinceEpoch() / trantor::Date::MICRO_SECONDS_PER_SEC; date.microSecondsSinceEpoch() / trantor::Date::MICRO_SECONDS_PER_SEC;
if (nowSecond == lastSecond) if (nowSecond == lastSecond)
@@ -1059,9 +1099,10 @@ const std::string &getHttpFullDateStr(const trantor::Date &date)
return lastTimeString; return lastTimeString;
} }
lastSecond = nowSecond; lastSecond = nowSecond;
dateToCustomFormattedString("%a, %d %b %Y %H:%M:%S GMT", lastTimeString.resize(128);
lastTimeString, int n = formatHttpDate(lastTimeString.data(), lastTimeString.size(), date);
date); n = std::clamp(n, 0, static_cast<int>(lastTimeString.size() - 1));
lastTimeString.resize(static_cast<size_t>(n));
return lastTimeString; return lastTimeString;
} }
+100
View File
@@ -0,0 +1,100 @@
/**
*
* @file ParsingUtils.h
* Shared parsing utilities for HTTP and multipart parsing
*
* Copyright 2024, Drogon. All rights reserved.
* https://github.com/drogonframework/drogon
* Use of this source code is governed by a MIT license
* that can be found in the License file.
*
* Drogon
*
*/
#pragma once
#include <cctype>
#include <string_view>
#include <utility>
namespace drogon
{
namespace utils
{
/**
* @brief Check if a string_view starts with another string_view
*/
inline bool startsWith(const std::string_view &a, const std::string_view &b)
{
if (a.size() < b.size())
{
return false;
}
for (size_t i = 0; i < b.size(); i++)
{
if (a[i] != b[i])
{
return false;
}
}
return true;
}
/**
* @brief Check if a string_view starts with another string_view
* (case-insensitive)
*/
inline bool startsWithIgnoreCase(const std::string_view &a,
const std::string_view &b)
{
if (a.size() < b.size())
{
return false;
}
for (size_t i = 0; i < b.size(); i++)
{
const auto lhs = std::tolower(static_cast<unsigned char>(a[i]));
const auto rhs = std::tolower(static_cast<unsigned char>(b[i]));
if (lhs != rhs)
{
return false;
}
}
return true;
}
/**
* @brief Parse a single HTTP header line into name and value
* @param begin Pointer to the start of the line
* @param end Pointer to the end of the line (not including CRLF)
* @return A pair of (header_name, header_value) string_views
*/
inline std::pair<std::string_view, std::string_view> parseLine(
const char *begin,
const char *end)
{
auto p = begin;
while (p != end)
{
if (*p == ':')
{
if (p + 1 != end && *(p + 1) == ' ')
{
return std::make_pair(std::string_view(begin, p - begin),
std::string_view(p + 2, end - p - 2));
}
else
{
return std::make_pair(std::string_view(begin, p - begin),
std::string_view(p + 1, end - p - 1));
}
}
++p;
}
return std::make_pair(std::string_view(), std::string_view());
}
} // namespace utils
} // namespace drogon
+1
View File
@@ -43,6 +43,7 @@ if(CMAKE_CXX_COMPILER_ID MATCHES "MSVC" AND BUILD_SHARED_LIBS)
else() else()
set(UNITTEST_SOURCES ${UNITTEST_SOURCES} ../src/HttpFileImpl.cc set(UNITTEST_SOURCES ${UNITTEST_SOURCES} ../src/HttpFileImpl.cc
unittests/HttpFileTest.cc unittests/HttpFileTest.cc
unittests/HttpMethodTest.cc
unittests/WebsocketResponseTest.cc) unittests/WebsocketResponseTest.cc)
endif() endif()
@@ -5,6 +5,7 @@
#include <string> #include <string>
#include <iostream> #include <iostream>
#include <fstream> #include <fstream>
#include <filesystem>
using namespace drogon; using namespace drogon;
@@ -100,23 +101,46 @@ DROGON_TEST(RequestStreamTest)
LOG_INFO << "Test request stream"; LOG_INFO << "Test request stream";
std::string filePath = "./中文.txt"; const auto uniqueSuffix = std::to_string(
std::ifstream file(filePath); std::chrono::steady_clock::now().time_since_epoch().count());
std::stringstream content; auto tempDir = std::make_shared<std::filesystem::path>(
REQUIRE(file.is_open()); std::filesystem::temp_directory_path() /
content << file.rdbuf(); ("request_stream_upload_test_" + uniqueSuffix));
std::filesystem::create_directories(*tempDir);
auto tempPath = std::make_shared<std::filesystem::path>(
*tempDir / std::filesystem::path(u8"中文.txt"));
tempPath->make_preferred();
{
std::ofstream out(*tempPath, std::ios::binary | std::ios::trunc);
REQUIRE(out.is_open());
out << "request-stream-upload-content\nline2\n";
}
req = HttpRequest::newFileUploadRequest({UploadFile{filePath}}); std::ifstream in(*tempPath, std::ios::binary);
REQUIRE(in.is_open());
std::stringstream ss;
ss << in.rdbuf();
const auto uploadContent = std::make_shared<std::string>(ss.str());
const auto uploadPathUtf8 = std::make_shared<std::string>([&tempPath]() {
auto u8Path = tempPath->u8string();
return std::string(reinterpret_cast<const char *>(u8Path.data()),
u8Path.size());
}());
req = HttpRequest::newFileUploadRequest({UploadFile{*uploadPathUtf8}});
req->setPath("/stream_upload_echo"); req->setPath("/stream_upload_echo");
req->setMethod(Post); req->setMethod(Post);
client->sendRequest(req, client->sendRequest(
[TEST_CTX, req,
content = content.str()](ReqResult r, [TEST_CTX, tempPath, tempDir, uploadPathUtf8, content = uploadContent](
const HttpResponsePtr &resp) { ReqResult r, const HttpResponsePtr &resp) {
CHECK(r == ReqResult::Ok); CHECK(r == ReqResult::Ok);
CHECK(resp->statusCode() == k200OK); CHECK(resp->statusCode() == k200OK);
CHECK(resp->body() == content); CHECK(resp->body() == *content);
}); std::error_code ec;
std::filesystem::remove(*tempPath, ec);
std::filesystem::remove(*tempDir, ec);
});
checkStreamRequest(TEST_CTX, checkStreamRequest(TEST_CTX,
client->getLoop(), client->getLoop(),
@@ -728,6 +728,23 @@ void doTest(const HttpClientPtr &client, std::shared_ptr<test::Case> TEST_CTX)
CHECK((*json)["P2"] == "test"); CHECK((*json)["P2"] == "test");
}); });
// Test file upload from memory
auto hello = std::make_shared<std::string>("hello world!");
UploadFile memfile(hello->data(),
hello->length(),
"hello_world.txt",
"hellofile",
ContentType::CT_TEXT_PLAIN);
req = HttpRequest::newFileUploadRequest({memfile});
req->setPath("/api/attachment/uploadMemory");
client->sendRequest(req,
[req, TEST_CTX, hello](ReqResult result,
const HttpResponsePtr &resp) {
REQUIRE(result == ReqResult::Ok);
REQUIRE(resp->contentType() == CT_TEXT_PLAIN);
CHECK(resp->getBody() == *hello);
});
// Test newFileResponse // Test newFileResponse
req = HttpRequest::newHttpRequest(); req = HttpRequest::newHttpRequest();
req->setPath("/RangeTestController/"); req->setPath("/RangeTestController/");
@@ -103,6 +103,32 @@ void Attachment::uploadImage(
callback(resp); callback(resp);
} }
void Attachment::uploadMemory(
const HttpRequestPtr &req,
std::function<void(const HttpResponsePtr &)> &&callback)
{
MultiPartParser fileUpload;
if (fileUpload.parse(req) == 0 && fileUpload.getFiles().size() == 1)
{
auto &file = fileUpload.getFiles()[0];
if (file.getItemName() == "hellofile")
{
auto resp = HttpResponse::newHttpResponse();
resp->setStatusCode(HttpStatusCode::k200OK);
resp->setContentTypeCode(ContentType::CT_TEXT_PLAIN);
std::string hello = std::string(file.fileData(), file.fileLength());
resp->setBody(std::move(hello));
callback(resp);
return;
}
}
LOG_DEBUG << "upload text from memory error!";
auto resp = HttpResponse::newHttpResponse();
resp->setStatusCode(HttpStatusCode::k400BadRequest);
callback(resp);
}
void Attachment::download( void Attachment::download(
const HttpRequestPtr &req, const HttpRequestPtr &req,
std::function<void(const HttpResponsePtr &)> &&callback) std::function<void(const HttpResponsePtr &)> &&callback)
@@ -12,6 +12,7 @@ class Attachment : public drogon::HttpController<Attachment>
METHOD_ADD(Attachment::get, "", Get); // Path is '/api/attachment' METHOD_ADD(Attachment::get, "", Get); // Path is '/api/attachment'
METHOD_ADD(Attachment::upload, "/upload", Post); METHOD_ADD(Attachment::upload, "/upload", Post);
METHOD_ADD(Attachment::uploadImage, "/uploadImage", Post); METHOD_ADD(Attachment::uploadImage, "/uploadImage", Post);
METHOD_ADD(Attachment::uploadMemory, "/uploadMemory", Post);
METHOD_ADD(Attachment::download, "/download", Get); METHOD_ADD(Attachment::download, "/download", Get);
METHOD_LIST_END METHOD_LIST_END
// your declaration of processing function maybe like this: // your declaration of processing function maybe like this:
@@ -21,6 +22,8 @@ class Attachment : public drogon::HttpController<Attachment>
std::function<void(const HttpResponsePtr &)> &&callback); std::function<void(const HttpResponsePtr &)> &&callback);
void uploadImage(const HttpRequestPtr &req, void uploadImage(const HttpRequestPtr &req,
std::function<void(const HttpResponsePtr &)> &&callback); std::function<void(const HttpResponsePtr &)> &&callback);
void uploadMemory(const HttpRequestPtr &req,
std::function<void(const HttpResponsePtr &)> &&callback);
void download(const HttpRequestPtr &req, void download(const HttpRequestPtr &req,
std::function<void(const HttpResponsePtr &)> &&callback); std::function<void(const HttpResponsePtr &)> &&callback);
}; };
@@ -9,6 +9,7 @@ DROGON_TEST(Base64)
auto decoded = drogon::utils::base64Decode(encoded); auto decoded = drogon::utils::base64Decode(encoded);
CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw=="); CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw==");
CHECK(decoded == in); CHECK(decoded == in);
CHECK(drogon::utils::isBase64(encoded));
SUBSECTION(InvalidChars) SUBSECTION(InvalidChars)
{ {
@@ -31,6 +32,7 @@ DROGON_TEST(Base64)
auto decoded = drogon::utils::base64Decode(encoded); auto decoded = drogon::utils::base64Decode(encoded);
CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw"); CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw");
CHECK(decoded == in); CHECK(decoded == in);
CHECK(drogon::utils::isBase64(encoded));
} }
SUBSECTION(LongString) SUBSECTION(LongString)
@@ -46,6 +48,9 @@ DROGON_TEST(Base64)
auto encoded = drogon::utils::base64Encode(in); auto encoded = drogon::utils::base64Encode(in);
auto decoded = drogon::utils::base64Decode(encoded); auto decoded = drogon::utils::base64Decode(encoded);
CHECK(decoded == in); CHECK(decoded == in);
CHECK(out == encoded);
CHECK(drogon::utils::isBase64(out));
CHECK(drogon::utils::isBase64(encoded));
} }
SUBSECTION(URLSafe) SUBSECTION(URLSafe)
@@ -55,6 +60,7 @@ DROGON_TEST(Base64)
auto decoded = drogon::utils::base64Decode(encoded); auto decoded = drogon::utils::base64Decode(encoded);
CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw=="); CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw==");
CHECK(decoded == in); CHECK(decoded == in);
CHECK(drogon::utils::isBase64(encoded));
} }
SUBSECTION(UnpaddedURLSafe) SUBSECTION(UnpaddedURLSafe)
@@ -64,6 +70,7 @@ DROGON_TEST(Base64)
auto decoded = drogon::utils::base64Decode(encoded); auto decoded = drogon::utils::base64Decode(encoded);
CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw"); CHECK(encoded == "ZHJvZ29uIGZyYW1ld29yaw");
CHECK(decoded == in); CHECK(decoded == in);
CHECK(drogon::utils::isBase64(encoded));
} }
SUBSECTION(LongURLSafe) SUBSECTION(LongURLSafe)
@@ -77,5 +84,24 @@ DROGON_TEST(Base64)
auto encoded = drogon::utils::base64Encode(in, true); auto encoded = drogon::utils::base64Encode(in, true);
auto decoded = drogon::utils::base64Decode(encoded); auto decoded = drogon::utils::base64Decode(encoded);
CHECK(decoded == in); CHECK(decoded == in);
CHECK(drogon::utils::isBase64(encoded));
}
SUBSECTION(emptyString)
{
auto encoded = "";
CHECK(!drogon::utils::isBase64(encoded));
}
SUBSECTION(size1Padding)
{
auto encoded = "ZHJvZ29uIGZyYW1ld29=";
CHECK(drogon::utils::isBase64(encoded));
}
SUBSECTION(size1PaddingNotModulo4)
{
auto encoded = "ZHJvZ29uIGZyYW1ld29ya=";
CHECK(!drogon::utils::isBase64(encoded));
} }
} }
@@ -66,3 +66,195 @@ DROGON_TEST(ResquestSetCustomContentTypeString)
req->setContentTypeString("thisdoesnotexist/unknown"); req->setContentTypeString("thisdoesnotexist/unknown");
CHECK(req->getContentType() == CT_CUSTOM); CHECK(req->getContentType() == CT_CUSTOM);
} }
DROGON_TEST(HttpOptionsHeadersResponse)
{
auto req = HttpRequest::newHttpRequest();
auto resp = HttpResponse::newOptionsResponse(req);
CHECK(!resp);
req->setMethod(HttpMethod::Options);
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Vary") == "Origin");
CHECK(resp->getHeader("Allow") == "OPTIONS");
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "");
CHECK(resp->getHeader("Access-Control-Allow-Methods") == "");
req->attributes()->insert("drogon.corsMethods",
std::string("GET, POST, OPTIONS"));
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Vary") == "Origin");
CHECK(resp->getHeader("Allow") == "GET, POST, OPTIONS");
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "");
CHECK(resp->getHeader("Access-Control-Allow-Methods") == "");
req->addHeader("Origin", "http://somepage");
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Vary") == "Origin");
CHECK(resp->getHeader("Allow") == "GET, POST, OPTIONS");
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "");
CHECK(resp->getHeader("Access-Control-Allow-Methods") == "");
}
DROGON_TEST(HttpCorsHeadersResponse)
{
auto req = HttpRequest::newHttpRequest();
req->addHeader("Origin", "");
req->addHeader("Access-Control-Request-Method", "OPTIONS");
auto resp = HttpResponse::newOptionsResponse(req);
CHECK(!resp);
// empty origin -> error
req->setMethod(HttpMethod::Options);
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k400BadRequest);
// null origin -> check if allowed or not
req->addHeader("Origin", "null");
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k403Forbidden);
resp = HttpResponse::newOptionsResponse(req, {}, true);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "null");
// normal origin but no requested method -> error
req->addHeader("Origin", "http://somepage");
req->addHeader("Access-Control-Request-Method", "");
resp = HttpResponse::newOptionsResponse(req, {}, true);
CHECK(resp->getStatusCode() == HttpStatusCode::k400BadRequest);
// valid CORS preflight request
req->addHeader("Access-Control-Request-Method", "OPTIONS");
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Vary") == "Origin");
CHECK(resp->getHeader("Allow") == "");
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "http://somepage");
CHECK(resp->getHeader("Access-Control-Allow-Methods") == "OPTIONS");
// origin validator
resp = HttpResponse::newOptionsResponse(req, [](std::string_view origin) {
return origin == "http://somepage";
});
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
resp = HttpResponse::newOptionsResponse(req, [](std::string_view origin) {
return origin != "http://somepage";
});
CHECK(resp->getStatusCode() == HttpStatusCode::k403Forbidden);
// unallowed method
req->addHeader("Access-Control-Request-Method", "PUT");
req->attributes()->insert("drogon.corsMethods",
std::string("GET,POST,OPTIONS"));
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k405MethodNotAllowed);
CHECK(resp->getHeader("Allow") == "GET,POST,OPTIONS");
CHECK(resp->getHeader("Access-Control-Allow-Methods") ==
"GET,POST,OPTIONS");
// allowed method
req->addHeader("Access-Control-Request-Method", "GET");
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Allow") == "");
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "http://somepage");
CHECK(resp->getHeader("Access-Control-Allow-Methods") ==
"GET,POST,OPTIONS");
CHECK(resp->getHeader("Access-Control-Allow-Credentials") == "");
CHECK(resp->getHeader("Access-Control-Allow-Private-Network") == "");
CHECK(resp->getHeader("Access-Control-Max-Age") == "");
// no restriction on requested headers
req->addHeader("Access-Control-Request-Headers", "X-Foo, X-Bar");
resp = HttpResponse::newOptionsResponse(req);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Allow-Headers") == "X-Bar,X-Foo");
// unallowed header
resp = HttpResponse::newOptionsResponse(req, {"X-Foo"});
CHECK(resp->getStatusCode() == HttpStatusCode::k403Forbidden);
CHECK(resp->getHeader("Access-Control-Allow-Headers") == "X-Foo");
// all requested headers allowed
resp = HttpResponse::newOptionsResponse(req, {"X-Foo", "X-Bar"});
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Allow-Headers") == "X-Bar,X-Foo");
// allow credentials
resp = HttpResponse::newOptionsResponse(req, nullptr, false, true);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Allow-Credentials") == "true");
// private network access
req->addHeader("Access-Control-Request-Private-Network", "true");
resp = HttpResponse::newOptionsResponse(req, nullptr, false, false, false);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Allow-Private-Network") == "");
resp = HttpResponse::newOptionsResponse(req, nullptr, false, false, true);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Allow-Private-Network") == "true");
// CORS max age
resp = HttpResponse::newOptionsResponse(
req, nullptr, false, false, false, 600);
CHECK(resp->getStatusCode() == HttpStatusCode::k204NoContent);
CHECK(resp->getHeader("Access-Control-Max-Age") == "600");
}
DROGON_TEST(AddHttpCorsHeaders)
{
using namespace std::literals;
// no Origin -> do nothing
auto req = HttpRequest::newHttpRequest();
req->setMethod(Get);
auto resp = HttpResponse::newHttpResponse();
resp->addCorsHeaders(req, {"X-Foo"}, true);
CHECK(resp->headers().empty());
// with Origin -> Allow-Origin + Vary (not overwritten) + Expose-Headers
req->addHeader("Origin", "http://somepage");
resp->addHeader("Vary", "X-SomeHeader");
resp->addCorsHeaders(req, {"X-Foo"});
CHECK(resp->getHeader("Vary") == "Origin,X-SomeHeader");
CHECK(resp->getHeader("Access-Control-Allow-Origin") == "http://somepage");
CHECK(resp->getHeader("Access-Control-Expose-Headers") == "X-Foo");
// add a new exposed header
resp->addCorsHeaders(req, {"X-Bar"});
CHECK(resp->getHeader("Access-Control-Expose-Headers") == "X-Bar,X-Foo");
// no duplicate Origin in Vary
CHECK(resp->getHeader("Vary") == "Origin,X-SomeHeader");
// check credentials (true/false/unchanged)
resp->addCorsHeaders(req, {}, true);
CHECK(resp->getHeader("Access-Control-Expose-Headers") == "X-Bar,X-Foo");
CHECK(resp->getHeader("Access-Control-Allow-Credentials") == "true");
resp->addCorsHeaders(req, {}, false);
CHECK(resp->getHeader("Access-Control-Allow-Credentials") == "");
resp->addCorsHeaders(req, {}, true);
resp->addCorsHeaders(req);
CHECK(resp->getHeader("Access-Control-Allow-Credentials") == "true");
}
DROGON_TEST(ClearHeaders)
{
auto req = HttpRequest::newHttpRequest();
// set a custom path to ensure it is not cleared
req->setPath("/api/test");
req->addHeader("X-Test", "value");
req->addHeader("Authorization", "Bearer token");
CHECK(req->headers().size() == 2);
CHECK(req->getHeader("X-Test") == "value");
CHECK(req->getHeader("Authorization") == "Bearer token");
req->clearHeaders();
CHECK(req->headers().empty());
// verify path unchanged
CHECK(req->path() == "/api/test");
}
@@ -0,0 +1,120 @@
#include <drogon/drogon_test.h>
#include <drogon/HttpTypes.h>
#include <trantor/utils/MsgBuffer.h>
#include "../../lib/src/HttpRequestImpl.h"
using namespace drogon;
// Helper: parse a method string through HttpRequestImpl::setMethod
static std::pair<bool, HttpMethod> parseMethod(const std::string &str)
{
HttpRequestImpl req(nullptr);
bool ok = req.setMethod(str.data(), str.data() + str.size());
return {ok, req.method()};
}
DROGON_TEST(StandardHttpMethods)
{
auto [ok, m] = parseMethod("GET");
CHECK(ok);
CHECK(m == Get);
std::tie(ok, m) = parseMethod("POST");
CHECK(ok);
CHECK(m == Post);
std::tie(ok, m) = parseMethod("PUT");
CHECK(ok);
CHECK(m == Put);
std::tie(ok, m) = parseMethod("DELETE");
CHECK(ok);
CHECK(m == Delete);
std::tie(ok, m) = parseMethod("HEAD");
CHECK(ok);
CHECK(m == Head);
std::tie(ok, m) = parseMethod("OPTIONS");
CHECK(ok);
CHECK(m == Options);
std::tie(ok, m) = parseMethod("PATCH");
CHECK(ok);
CHECK(m == Patch);
}
DROGON_TEST(WebDavMethods)
{
auto [ok, m] = parseMethod("PROPFIND");
CHECK(ok);
CHECK(m == Propfind);
std::tie(ok, m) = parseMethod("MKCOL");
CHECK(ok);
CHECK(m == Mkcol);
std::tie(ok, m) = parseMethod("COPY");
CHECK(ok);
CHECK(m == Copy);
std::tie(ok, m) = parseMethod("MOVE");
CHECK(ok);
CHECK(m == Move);
}
DROGON_TEST(WebDavMethodStrings)
{
CHECK(to_string_view(Propfind) == "PROPFIND");
CHECK(to_string_view(Mkcol) == "MKCOL");
CHECK(to_string_view(Copy) == "COPY");
CHECK(to_string_view(Move) == "MOVE");
}
// Helper: serialize a request and return the first line (method + path)
static std::string serializeMethod(HttpMethod method)
{
HttpRequestImpl req(nullptr);
req.setMethod(method);
req.setPath("/test");
trantor::MsgBuffer buf;
req.appendToBuffer(&buf);
std::string result(buf.peek(), buf.readableBytes());
// Return just up to the first space after the method
auto pos = result.find(' ');
return result.substr(0, pos);
}
DROGON_TEST(MethodSerialization)
{
CHECK(serializeMethod(Get) == "GET");
CHECK(serializeMethod(Post) == "POST");
CHECK(serializeMethod(Put) == "PUT");
CHECK(serializeMethod(Delete) == "DELETE");
CHECK(serializeMethod(Head) == "HEAD");
CHECK(serializeMethod(Options) == "OPTIONS");
CHECK(serializeMethod(Patch) == "PATCH");
CHECK(serializeMethod(Propfind) == "PROPFIND");
CHECK(serializeMethod(Mkcol) == "MKCOL");
CHECK(serializeMethod(Copy) == "COPY");
CHECK(serializeMethod(Move) == "MOVE");
}
DROGON_TEST(InvalidMethodsRejected)
{
auto [ok, m] = parseMethod("INVALID");
CHECK(!ok);
CHECK(m == Invalid);
std::tie(ok, m) = parseMethod("LOCK");
CHECK(!ok);
CHECK(m == Invalid);
std::tie(ok, m) = parseMethod("");
CHECK(!ok);
CHECK(m == Invalid);
std::tie(ok, m) = parseMethod("G");
CHECK(!ok);
CHECK(m == Invalid);
}
@@ -67,6 +67,69 @@ struct Filter
std::string value; std::string value;
}; };
/**
* @brief Represents a SQL JOIN clause.
*/
enum class JoinType
{
InnerJoin,
LeftJoin,
RightJoin,
FullJoin
};
inline std::string to_join_string(JoinType type)
{
switch (type)
{
case JoinType::InnerJoin:
return "INNER JOIN";
case JoinType::LeftJoin:
return "LEFT JOIN";
case JoinType::RightJoin:
return "RIGHT JOIN";
case JoinType::FullJoin:
return "FULL JOIN";
}
// Should never reach here
return "INNER JOIN";
}
struct JoinClause
{
JoinType type;
std::string table;
std::string onLeft; // e.g. "users.id"
std::string onRight; // e.g. "posts.user_id"
};
/**
* @brief Validate that a string is a safe SQL identifier.
*
* Only allows alphanumeric characters, underscores, and dots
* (for table.column notation). This prevents SQL injection when
* building JOIN clauses from user-provided identifiers.
*
* @param identifier The identifier to validate.
* @return true if the identifier is safe to use in SQL.
*/
inline bool isValidSqlIdentifier(const std::string &identifier)
{
if (identifier.empty())
{
return false;
}
for (auto c : identifier)
{
if (!std::isalnum(static_cast<unsigned char>(c)) && c != '_' &&
c != '.')
{
return false;
}
}
return true;
}
// Forward declaration to be a friend // Forward declaration to be a friend
template <typename T, bool SelectAll, bool Single = false> template <typename T, bool SelectAll, bool Single = false>
class TransformBuilder; class TransformBuilder;
@@ -87,6 +150,7 @@ class BaseBuilder
std::string from_; std::string from_;
std::string columns_; std::string columns_;
std::vector<Filter> filters_; std::vector<Filter> filters_;
std::vector<JoinClause> joins_;
std::optional<std::uint64_t> limit_; std::optional<std::uint64_t> limit_;
std::optional<std::uint64_t> offset_; std::optional<std::uint64_t> offset_;
// The order is important; use vector<pair> instead of unordered_map and // The order is important; use vector<pair> instead of unordered_map and
@@ -122,6 +186,11 @@ class BaseBuilder
}; };
std::string sql = "select " + columns_ + " from " + from_; std::string sql = "select " + columns_ + " from " + from_;
for (const auto &join : joins_)
{
sql += " " + to_join_string(join.type) + " " + join.table + " ON " +
join.onLeft + " = " + join.onRight;
}
if (!filters_.empty()) if (!filters_.empty())
{ {
sql += " where " + filters_[0].column + " " + sql += " where " + filters_[0].column + " " +
@@ -211,6 +211,54 @@ class CoroMapper : public Mapper<T>
return *this; return *this;
} }
/**
* @brief Add an INNER JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of ON (e.g. "users.id").
* @param onRight The right side of ON (e.g. "posts.user_id").
* @return CoroMapper<T>& The CoroMapper itself.
*/
CoroMapper<T> &innerJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
Mapper<T>::innerJoin(table, onLeft, onRight);
return *this;
}
/**
* @brief Add a LEFT JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of ON (e.g. "users.id").
* @param onRight The right side of ON (e.g. "posts.user_id").
* @return CoroMapper<T>& The CoroMapper itself.
*/
CoroMapper<T> &leftJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
Mapper<T>::leftJoin(table, onLeft, onRight);
return *this;
}
/**
* @brief Add a RIGHT JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of ON (e.g. "users.id").
* @param onRight The right side of ON (e.g. "posts.user_id").
* @return CoroMapper<T>& The CoroMapper itself.
*/
CoroMapper<T> &rightJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
Mapper<T>::rightJoin(table, onLeft, onRight);
return *this;
}
// Read api for coroutines // Read api for coroutines
inline internal::MapperAwaiter<std::vector<T>> findAll() inline internal::MapperAwaiter<std::vector<T>> findAll()
@@ -225,6 +273,7 @@ class CoroMapper : public Mapper<T>
ExceptPtrCallback &&errCallback) { ExceptPtrCallback &&errCallback) {
std::string sql = "select count(*) from "; std::string sql = "select count(*) from ";
sql += T::tableName; sql += T::tableName;
sql += this->joinString_;
if (criteria) if (criteria)
{ {
sql += " where "; sql += " where ";
@@ -250,6 +299,7 @@ class CoroMapper : public Mapper<T>
ExceptPtrCallback &&errCallback) { ExceptPtrCallback &&errCallback) {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += this->joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -311,6 +361,7 @@ class CoroMapper : public Mapper<T>
ExceptPtrCallback &&errCallback) { ExceptPtrCallback &&errCallback) {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += this->joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
+41 -6
View File
@@ -43,6 +43,15 @@ using ExceptionCallback = std::function<void(const DrogonDbException &)>;
class Transaction; class Transaction;
class DbClient; class DbClient;
/// Transaction locking mode.
enum class TransactionType
{
Deferred, ///< BEGIN — lock acquired on first write (default)
Immediate, ///< BEGIN IMMEDIATE — write lock acquired upfront (SQLite only)
Exclusive, ///< BEGIN EXCLUSIVE — exclusive lock acquired upfront (SQLite
///< only)
};
namespace internal namespace internal
{ {
#ifdef __cpp_impl_coroutine #ifdef __cpp_impl_coroutine
@@ -73,7 +82,10 @@ struct [[nodiscard]] SqlAwaiter : public CallbackAwaiter<Result>
struct [[nodiscard]] TransactionAwaiter struct [[nodiscard]] TransactionAwaiter
: public CallbackAwaiter<std::shared_ptr<Transaction> > : public CallbackAwaiter<std::shared_ptr<Transaction> >
{ {
explicit TransactionAwaiter(DbClient *client) : client_(client) explicit TransactionAwaiter(
DbClient *client,
TransactionType transType = TransactionType::Deferred)
: client_(client), transType_(transType)
{ {
} }
@@ -81,6 +93,7 @@ struct [[nodiscard]] TransactionAwaiter
private: private:
DbClient *client_; DbClient *client_;
TransactionType transType_;
}; };
#endif #endif
@@ -269,7 +282,16 @@ class DROGON_EXPORT DbClient : public trantor::NonCopyable
*/ */
virtual std::shared_ptr<Transaction> newTransaction( virtual std::shared_ptr<Transaction> newTransaction(
const std::function<void(bool)> &commitCallback = const std::function<void(bool)> &commitCallback =
std::function<void(bool)>()) noexcept(false) = 0; std::function<void(bool)>(),
TransactionType transType =
TransactionType::Deferred) noexcept(false) = 0;
/// Convenience overload: create a transaction with a specific locking mode.
std::shared_ptr<Transaction> newTransaction(
TransactionType transType) noexcept(false)
{
return newTransaction(std::function<void(bool)>(), transType);
}
/// Create a transaction object in asynchronous mode. /// Create a transaction object in asynchronous mode.
/** /**
@@ -278,12 +300,24 @@ class DROGON_EXPORT DbClient : public trantor::NonCopyable
*/ */
virtual void newTransactionAsync( virtual void newTransactionAsync(
const std::function<void(const std::shared_ptr<Transaction> &)> const std::function<void(const std::shared_ptr<Transaction> &)>
&callback) = 0; &callback,
TransactionType transType = TransactionType::Deferred) = 0;
/// Convenience overload: create an async transaction with a specific
/// locking mode, with transType as the first argument.
void newTransactionAsync(
TransactionType transType,
const std::function<void(const std::shared_ptr<Transaction> &)>
&callback)
{
newTransactionAsync(callback, transType);
}
#ifdef __cpp_impl_coroutine #ifdef __cpp_impl_coroutine
orm::internal::TransactionAwaiter newTransactionCoro() orm::internal::TransactionAwaiter newTransactionCoro(
TransactionType transType = TransactionType::Deferred)
{ {
return orm::internal::TransactionAwaiter(this); return orm::internal::TransactionAwaiter(this, transType);
} }
#endif #endif
@@ -408,7 +442,8 @@ inline void internal::TransactionAwaiter::await_suspend(
else else
setValue(transaction); setValue(transaction);
handle.resume(); handle.resume();
}); },
transType_);
} }
#endif #endif
@@ -157,6 +157,90 @@ class FilterBuilder : public TransformBuilder<T, SelectAll, false>
this->filters_.push_back({column, CompareOperator::Like, pattern}); this->filters_.push_back({column, CompareOperator::Like, pattern});
return *this; return *this;
} }
/**
* @brief Add an INNER JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of the ON condition (e.g. "users.id").
* @param onRight The right side of the ON condition (e.g.
* "posts.user_id").
*
* @return FilterBuilder& The FilterBuilder itself.
*/
inline FilterBuilder &innerJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
this->joins_.push_back({JoinType::InnerJoin, table, onLeft, onRight});
return *this;
}
/**
* @brief Add a LEFT JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of the ON condition (e.g. "users.id").
* @param onRight The right side of the ON condition (e.g.
* "posts.user_id").
*
* @return FilterBuilder& The FilterBuilder itself.
*/
inline FilterBuilder &leftJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
this->joins_.push_back({JoinType::LeftJoin, table, onLeft, onRight});
return *this;
}
/**
* @brief Add a RIGHT JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of the ON condition (e.g. "users.id").
* @param onRight The right side of the ON condition (e.g.
* "posts.user_id").
*
* @return FilterBuilder& The FilterBuilder itself.
*/
inline FilterBuilder &rightJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
this->joins_.push_back({JoinType::RightJoin, table, onLeft, onRight});
return *this;
}
/**
* @brief Add a FULL JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of the ON condition (e.g. "users.id").
* @param onRight The right side of the ON condition (e.g.
* "posts.user_id").
*
* @return FilterBuilder& The FilterBuilder itself.
*/
inline FilterBuilder &fullJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
this->joins_.push_back({JoinType::FullJoin, table, onLeft, onRight});
return *this;
}
}; };
} // namespace orm } // namespace orm
} // namespace drogon } // namespace drogon
+90 -3
View File
@@ -14,6 +14,7 @@
#pragma once #pragma once
#include <drogon/orm/Criteria.h> #include <drogon/orm/Criteria.h>
#include <drogon/orm/BaseBuilder.h>
#include <drogon/orm/DbClient.h> #include <drogon/orm/DbClient.h>
#include <drogon/utils/Utilities.h> #include <drogon/utils/Utilities.h>
#include <string> #include <string>
@@ -178,6 +179,78 @@ class Mapper
*/ */
Mapper<T> &forUpdate(); Mapper<T> &forUpdate();
/**
* @brief Add an INNER JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of ON (e.g. "users.id").
* @param onRight The right side of ON (e.g. "posts.user_id").
* @return Mapper<T>& The Mapper itself.
*/
Mapper<T> &innerJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
joinString_ += " INNER JOIN ";
joinString_ += table;
joinString_ += " ON ";
joinString_ += onLeft;
joinString_ += " = ";
joinString_ += onRight;
return *this;
}
/**
* @brief Add a LEFT JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of ON (e.g. "users.id").
* @param onRight The right side of ON (e.g. "posts.user_id").
* @return Mapper<T>& The Mapper itself.
*/
Mapper<T> &leftJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
joinString_ += " LEFT JOIN ";
joinString_ += table;
joinString_ += " ON ";
joinString_ += onLeft;
joinString_ += " = ";
joinString_ += onRight;
return *this;
}
/**
* @brief Add a RIGHT JOIN clause to the query.
*
* @param table The table to join.
* @param onLeft The left side of ON (e.g. "users.id").
* @param onRight The right side of ON (e.g. "posts.user_id").
* @return Mapper<T>& The Mapper itself.
*/
Mapper<T> &rightJoin(const std::string &table,
const std::string &onLeft,
const std::string &onRight)
{
assert(isValidSqlIdentifier(table));
assert(isValidSqlIdentifier(onLeft));
assert(isValidSqlIdentifier(onRight));
joinString_ += " RIGHT JOIN ";
joinString_ += table;
joinString_ += " ON ";
joinString_ += onLeft;
joinString_ += " = ";
joinString_ += onRight;
return *this;
}
using SingleRowCallback = std::function<void(T)>; using SingleRowCallback = std::function<void(T)>;
using MultipleRowsCallback = std::function<void(std::vector<T>)>; using MultipleRowsCallback = std::function<void(std::vector<T>)>;
using CountCallback = std::function<void(const size_t)>; using CountCallback = std::function<void(const size_t)>;
@@ -719,6 +792,7 @@ class Mapper
size_t limit_{0}; size_t limit_{0};
size_t offset_{0}; size_t offset_{0};
std::string orderByString_; std::string orderByString_;
std::string joinString_;
bool forUpdate_{false}; bool forUpdate_{false};
void clear() void clear()
@@ -726,6 +800,7 @@ class Mapper
limit_ = 0; limit_ = 0;
offset_ = 0; offset_ = 0;
orderByString_.clear(); orderByString_.clear();
joinString_.clear();
forUpdate_ = false; forUpdate_ = false;
} }
@@ -792,6 +867,7 @@ inline T Mapper<T>::findOne(const Criteria &criteria) noexcept(false)
{ {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -849,6 +925,7 @@ inline void Mapper<T>::findOne(const Criteria &criteria,
{ {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -904,6 +981,7 @@ inline std::future<T> Mapper<T>::findFutureOne(
{ {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -964,6 +1042,7 @@ inline std::vector<T> Mapper<T>::findBy(const Criteria &criteria) noexcept(
{ {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -1003,9 +1082,10 @@ inline std::vector<T> Mapper<T>::findBy(const Criteria &criteria) noexcept(
binder.exec(); // exec may be throw exception; binder.exec(); // exec may be throw exception;
} }
std::vector<T> ret; std::vector<T> ret;
ret.reserve(r.size());
for (auto const &row : r) for (auto const &row : r)
{ {
ret.push_back(T(row)); ret.emplace_back(row);
} }
return ret; return ret;
} }
@@ -1017,6 +1097,7 @@ inline void Mapper<T>::findBy(const Criteria &criteria,
{ {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -1051,6 +1132,7 @@ inline void Mapper<T>::findBy(const Criteria &criteria,
clear(); clear();
binder >> [rcb](const Result &r) { binder >> [rcb](const Result &r) {
std::vector<T> ret; std::vector<T> ret;
ret.reserve(r.size());
for (auto const &row : r) for (auto const &row : r)
{ {
ret.emplace_back(row); ret.emplace_back(row);
@@ -1066,6 +1148,7 @@ inline std::future<std::vector<T>> Mapper<T>::findFutureBy(
{ {
std::string sql = "select * from "; std::string sql = "select * from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
bool hasParameters = false; bool hasParameters = false;
if (criteria) if (criteria)
{ {
@@ -1102,11 +1185,12 @@ inline std::future<std::vector<T>> Mapper<T>::findFutureBy(
std::make_shared<std::promise<std::vector<T>>>(); std::make_shared<std::promise<std::vector<T>>>();
binder >> [prom](const Result &r) { binder >> [prom](const Result &r) {
std::vector<T> ret; std::vector<T> ret;
ret.reserve(r.size());
for (auto const &row : r) for (auto const &row : r)
{ {
ret.push_back(T(row)); ret.emplace_back(row);
} }
prom->set_value(ret); prom->set_value(std::move(ret));
}; };
binder >> [prom](const std::exception_ptr &e) { prom->set_exception(e); }; binder >> [prom](const std::exception_ptr &e) { prom->set_exception(e); };
binder.exec(); binder.exec();
@@ -1137,6 +1221,7 @@ inline size_t Mapper<T>::count(const Criteria &criteria) noexcept(false)
{ {
std::string sql = "select count(*) from "; std::string sql = "select count(*) from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
if (criteria) if (criteria)
{ {
sql += " where "; sql += " where ";
@@ -1164,6 +1249,7 @@ inline void Mapper<T>::count(const Criteria &criteria,
{ {
std::string sql = "select count(*) from "; std::string sql = "select count(*) from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
if (criteria) if (criteria)
{ {
sql += " where "; sql += " where ";
@@ -1187,6 +1273,7 @@ inline std::future<size_t> Mapper<T>::countFuture(
{ {
std::string sql = "select count(*) from "; std::string sql = "select count(*) from ";
sql += T::tableName; sql += T::tableName;
sql += joinString_;
if (criteria) if (criteria)
{ {
sql += " where "; sql += " where ";
+26 -13
View File
@@ -197,7 +197,8 @@ void DbClientImpl::execSql(
} }
void DbClientImpl::newTransactionAsync( void DbClientImpl::newTransactionAsync(
const std::function<void(const std::shared_ptr<Transaction> &)> &callback) const std::function<void(const std::shared_ptr<Transaction> &)> &callback,
TransactionType transType)
{ {
DbConnectionPtr conn; DbConnectionPtr conn;
{ {
@@ -231,7 +232,7 @@ void DbClientImpl::newTransactionAsync(
iter != transCallbacks_.end(); iter != transCallbacks_.end();
++iter) ++iter)
{ {
if (cbPtr == *iter) if (cbPtr == iter->first)
{ {
transCallbacks_.erase(iter); transCallbacks_.erase(iter);
break; break;
@@ -251,24 +252,29 @@ void DbClientImpl::newTransactionAsync(
(*newCallbackPtr) = callbackPtr; (*newCallbackPtr) = callbackPtr;
timeoutFlagPtr->runTimer(); timeoutFlagPtr->runTimer();
} }
transCallbacks_.push_back(callbackPtr); transCallbacks_.push_back({callbackPtr, transType});
} }
} }
if (conn) if (conn)
{ {
makeTrans(conn, makeTrans(conn,
std::function<void(const std::shared_ptr<Transaction> &)>( std::function<void(const std::shared_ptr<Transaction> &)>(
callback)); callback),
transType);
} }
} }
void DbClientImpl::makeTrans( void DbClientImpl::makeTrans(
const DbConnectionPtr &conn, const DbConnectionPtr &conn,
std::function<void(const std::shared_ptr<Transaction> &)> &&callback) std::function<void(const std::shared_ptr<Transaction> &)> &&callback,
TransactionType transType)
{ {
std::weak_ptr<DbClientImpl> weakThis = shared_from_this(); std::weak_ptr<DbClientImpl> weakThis = shared_from_this();
auto trans = std::make_shared<TransactionImpl>( auto trans = std::make_shared<TransactionImpl>(
type_, conn, std::function<void(bool)>(), [weakThis, conn]() { type_,
conn,
std::function<void(bool)>(),
[weakThis, conn]() {
auto thisPtr = weakThis.lock(); auto thisPtr = weakThis.lock();
if (!thisPtr) if (!thisPtr)
return; return;
@@ -306,7 +312,8 @@ void DbClientImpl::makeTrans(
}); });
thisPtr->handleNewTask(conn); thisPtr->handleNewTask(conn);
}); });
}); },
transType);
trans->doBegin(); trans->doBegin();
if (timeout_ > 0.0) if (timeout_ > 0.0)
{ {
@@ -317,13 +324,16 @@ void DbClientImpl::makeTrans(
} }
std::shared_ptr<Transaction> DbClientImpl::newTransaction( std::shared_ptr<Transaction> DbClientImpl::newTransaction(
const std::function<void(bool)> &commitCallback) noexcept(false) const std::function<void(bool)> &commitCallback,
TransactionType transType) noexcept(false)
{ {
std::promise<std::shared_ptr<Transaction>> pro; std::promise<std::shared_ptr<Transaction>> pro;
auto f = pro.get_future(); auto f = pro.get_future();
newTransactionAsync([&pro](const std::shared_ptr<Transaction> &trans) { newTransactionAsync(
pro.set_value(trans); [&pro](const std::shared_ptr<Transaction> &trans) {
}); pro.set_value(trans);
},
transType);
auto trans = f.get(); auto trans = f.get();
if (!trans) if (!trans)
{ {
@@ -336,12 +346,15 @@ std::shared_ptr<Transaction> DbClientImpl::newTransaction(
void DbClientImpl::handleNewTask(const DbConnectionPtr &connPtr) void DbClientImpl::handleNewTask(const DbConnectionPtr &connPtr)
{ {
std::function<void(const std::shared_ptr<Transaction> &)> transCallback; std::function<void(const std::shared_ptr<Transaction> &)> transCallback;
TransactionType transType{TransactionType::Deferred};
std::shared_ptr<SqlCmd> cmd; std::shared_ptr<SqlCmd> cmd;
{ {
std::lock_guard<std::mutex> guard(connectionsMutex_); std::lock_guard<std::mutex> guard(connectionsMutex_);
if (!transCallbacks_.empty()) if (!transCallbacks_.empty())
{ {
transCallback = std::move(*(transCallbacks_.front())); auto &entry = transCallbacks_.front();
transCallback = std::move(*entry.first);
transType = entry.second;
transCallbacks_.pop_front(); transCallbacks_.pop_front();
} }
else if (!sqlCmdBuffer_.empty()) else if (!sqlCmdBuffer_.empty())
@@ -358,7 +371,7 @@ void DbClientImpl::handleNewTask(const DbConnectionPtr &connPtr)
} }
if (transCallback) if (transCallback)
{ {
makeTrans(connPtr, std::move(transCallback)); makeTrans(connPtr, std::move(transCallback), transType);
return; return;
} }
if (cmd) if (cmd)
+12 -6
View File
@@ -52,10 +52,13 @@ class DbClientImpl : public DbClient,
&&exceptCallback) override; &&exceptCallback) override;
std::shared_ptr<Transaction> newTransaction( std::shared_ptr<Transaction> newTransaction(
const std::function<void(bool)> &commitCallback = const std::function<void(bool)> &commitCallback =
std::function<void(bool)>()) noexcept(false) override; std::function<void(bool)>(),
TransactionType transType =
TransactionType::Deferred) noexcept(false) override;
void newTransactionAsync( void newTransactionAsync(
const std::function<void(const std::shared_ptr<Transaction> &)> const std::function<void(const std::shared_ptr<Transaction> &)>
&callback) override; &callback,
TransactionType transType = TransactionType::Deferred) override;
bool hasAvailableConnections() const noexcept override; bool hasAvailableConnections() const noexcept override;
void setTimeout(double timeout) override void setTimeout(double timeout) override
@@ -78,16 +81,19 @@ class DbClientImpl : public DbClient,
void makeTrans( void makeTrans(
const DbConnectionPtr &conn, const DbConnectionPtr &conn,
std::function<void(const std::shared_ptr<Transaction> &)> &&callback); std::function<void(const std::shared_ptr<Transaction> &)> &&callback,
TransactionType transType = TransactionType::Deferred);
mutable std::mutex connectionsMutex_; mutable std::mutex connectionsMutex_;
std::unordered_set<DbConnectionPtr> connections_; std::unordered_set<DbConnectionPtr> connections_;
std::unordered_set<DbConnectionPtr> readyConnections_; std::unordered_set<DbConnectionPtr> readyConnections_;
std::unordered_set<DbConnectionPtr> busyConnections_; std::unordered_set<DbConnectionPtr> busyConnections_;
std::list<std::shared_ptr< using TransCallbackEntry =
std::function<void(const std::shared_ptr<Transaction> &)>>> std::pair<std::shared_ptr<std::function<void(
transCallbacks_; const std::shared_ptr<Transaction> &)>>,
TransactionType>;
std::list<TransCallbackEntry> transCallbacks_;
std::deque<std::shared_ptr<SqlCmd>> sqlCmdBuffer_; std::deque<std::shared_ptr<SqlCmd>> sqlCmdBuffer_;
+24 -12
View File
@@ -230,7 +230,8 @@ void DbClientLockFree::execSql(
} }
std::shared_ptr<Transaction> DbClientLockFree::newTransaction( std::shared_ptr<Transaction> DbClientLockFree::newTransaction(
const std::function<void(bool)> &) noexcept(false) const std::function<void(bool)> &,
TransactionType) noexcept(false)
{ {
// Don't support transaction; // Don't support transaction;
LOG_ERROR LOG_ERROR
@@ -241,7 +242,8 @@ std::shared_ptr<Transaction> DbClientLockFree::newTransaction(
} }
void DbClientLockFree::newTransactionAsync( void DbClientLockFree::newTransactionAsync(
const std::function<void(const std::shared_ptr<Transaction> &)> &callback) const std::function<void(const std::shared_ptr<Transaction> &)> &callback,
TransactionType transType)
{ {
loop_->assertInLoopThread(); loop_->assertInLoopThread();
for (auto &conn : connections_) for (auto &conn : connections_)
@@ -250,7 +252,8 @@ void DbClientLockFree::newTransactionAsync(
{ {
makeTrans(conn, makeTrans(conn,
std::function<void(const std::shared_ptr<Transaction> &)>( std::function<void(const std::shared_ptr<Transaction> &)>(
callback)); callback),
transType);
return; return;
} }
} }
@@ -272,7 +275,7 @@ void DbClientLockFree::newTransactionAsync(
iter != transCallbacks_.end(); iter != transCallbacks_.end();
++iter) ++iter)
{ {
if (cbPtr == *iter) if (cbPtr == iter->first)
{ {
transCallbacks_.erase(iter); transCallbacks_.erase(iter);
break; break;
@@ -292,16 +295,20 @@ void DbClientLockFree::newTransactionAsync(
*newCallbackPtr = callbackPtr; *newCallbackPtr = callbackPtr;
timeoutFlagPtr->runTimer(); timeoutFlagPtr->runTimer();
} }
transCallbacks_.push_back(callbackPtr); transCallbacks_.push_back({callbackPtr, transType});
} }
void DbClientLockFree::makeTrans( void DbClientLockFree::makeTrans(
const DbConnectionPtr &conn, const DbConnectionPtr &conn,
std::function<void(const std::shared_ptr<Transaction> &)> &&callback) std::function<void(const std::shared_ptr<Transaction> &)> &&callback,
TransactionType transType)
{ {
std::weak_ptr<DbClientLockFree> weakThis = shared_from_this(); std::weak_ptr<DbClientLockFree> weakThis = shared_from_this();
auto trans = std::make_shared<TransactionImpl>( auto trans = std::make_shared<TransactionImpl>(
type_, conn, std::function<void(bool)>(), [weakThis, conn]() { type_,
conn,
std::function<void(bool)>(),
[weakThis, conn]() {
auto thisPtr = weakThis.lock(); auto thisPtr = weakThis.lock();
if (!thisPtr) if (!thisPtr)
return; return;
@@ -312,9 +319,11 @@ void DbClientLockFree::makeTrans(
} }
if (!thisPtr->transCallbacks_.empty()) if (!thisPtr->transCallbacks_.empty())
{ {
auto callback = std::move(thisPtr->transCallbacks_.front()); auto &entry = thisPtr->transCallbacks_.front();
auto nextCallback = std::move(*entry.first);
auto nextType = entry.second;
thisPtr->transCallbacks_.pop_front(); thisPtr->transCallbacks_.pop_front();
thisPtr->makeTrans(conn, std::move(*callback)); thisPtr->makeTrans(conn, std::move(nextCallback), nextType);
return; return;
} }
@@ -342,7 +351,8 @@ void DbClientLockFree::makeTrans(
break; break;
} }
} }
}); },
transType);
transSet_.insert(conn); transSet_.insert(conn);
trans->doBegin(); trans->doBegin();
if (timeout_ > 0.0) if (timeout_ > 0.0)
@@ -360,9 +370,11 @@ void DbClientLockFree::handleNewTask(const DbConnectionPtr &conn)
if (!transCallbacks_.empty()) if (!transCallbacks_.empty())
{ {
auto callback = std::move(transCallbacks_.front()); auto &entry = transCallbacks_.front();
auto callback = std::move(*entry.first);
auto transType = entry.second;
transCallbacks_.pop_front(); transCallbacks_.pop_front();
makeTrans(conn, std::move(*callback)); makeTrans(conn, std::move(callback), transType);
return; return;
} }
+12 -6
View File
@@ -55,10 +55,13 @@ class DbClientLockFree : public DbClient,
&&exceptCallback) override; &&exceptCallback) override;
std::shared_ptr<Transaction> newTransaction( std::shared_ptr<Transaction> newTransaction(
const std::function<void(bool)> &commitCallback = const std::function<void(bool)> &commitCallback =
std::function<void(bool)>()) noexcept(false) override; std::function<void(bool)>(),
TransactionType transType =
TransactionType::Deferred) noexcept(false) override;
void newTransactionAsync( void newTransactionAsync(
const std::function<void(const std::shared_ptr<Transaction> &)> const std::function<void(const std::shared_ptr<Transaction> &)>
&callback) override; &callback,
TransactionType transType = TransactionType::Deferred) override;
bool hasAvailableConnections() const noexcept override; bool hasAvailableConnections() const noexcept override;
void setTimeout(double timeout) override void setTimeout(double timeout) override
@@ -78,15 +81,18 @@ class DbClientLockFree : public DbClient,
std::unordered_set<DbConnectionPtr> transSet_; std::unordered_set<DbConnectionPtr> transSet_;
std::deque<std::shared_ptr<SqlCmd>> sqlCmdBuffer_; std::deque<std::shared_ptr<SqlCmd>> sqlCmdBuffer_;
std::list<std::shared_ptr< using TransCallbackEntry =
std::function<void(const std::shared_ptr<Transaction> &)>>> std::pair<std::shared_ptr<std::function<void(
transCallbacks_; const std::shared_ptr<Transaction> &)>>,
TransactionType>;
std::list<TransCallbackEntry> transCallbacks_;
double timeout_{-1.0}; double timeout_{-1.0};
void makeTrans( void makeTrans(
const DbConnectionPtr &conn, const DbConnectionPtr &conn,
std::function<void(const std::shared_ptr<Transaction> &)> &&callback); std::function<void(const std::shared_ptr<Transaction> &)> &&callback,
TransactionType transType = TransactionType::Deferred);
void execSqlWithTimeout( void execSqlWithTimeout(
const char *sql, const char *sql,
size_t sqlLength, size_t sqlLength,
+53 -19
View File
@@ -23,11 +23,13 @@ using namespace drogon;
TransactionImpl::TransactionImpl(ClientType type, TransactionImpl::TransactionImpl(ClientType type,
const DbConnectionPtr &connPtr, const DbConnectionPtr &connPtr,
std::function<void(bool)> commitCallback, std::function<void(bool)> commitCallback,
std::function<void()> usedUpCallback) std::function<void()> usedUpCallback,
TransactionType transType)
: connectionPtr_(connPtr), : connectionPtr_(connPtr),
usedUpCallback_(std::move(usedUpCallback)), usedUpCallback_(std::move(usedUpCallback)),
loop_(connPtr->loop()), loop_(connPtr->loop()),
commitCallback_(std::move(commitCallback)) commitCallback_(std::move(commitCallback)),
transactionType_(transType)
{ {
type_ = type; type_ = type;
} }
@@ -203,6 +205,8 @@ void TransactionImpl::execNewTask()
{ {
loop_->assertInLoopThread(); loop_->assertInLoopThread();
thisPtr_.reset(); thisPtr_.reset();
if (!isWorking_)
return;
assert(isWorking_); assert(isWorking_);
if (!isCommitedOrRolledback_) if (!isCommitedOrRolledback_)
{ {
@@ -244,27 +248,51 @@ void TransactionImpl::execNewTask()
else else
{ {
isWorking_ = false; isWorking_ = false;
if (!sqlCmdBuffer_.empty()) failBufferedCommands(std::make_exception_ptr(
TransactionRollback("The transaction has been rolled back")));
releaseConnection();
}
}
void TransactionImpl::releaseConnection()
{
if (usedUpCallback_)
{
usedUpCallback_();
usedUpCallback_ = std::function<void()>();
}
}
void TransactionImpl::failBufferedCommands(const std::exception_ptr &ePtr)
{
std::list<SqlCmdPtr> pendingCmds;
pendingCmds.swap(sqlCmdBuffer_);
for (auto &cmd : pendingCmds)
{
cmd->thisPtr_.reset();
if (cmd->exceptionCallback_)
{ {
auto exceptPtr = std::make_exception_ptr( cmd->exceptionCallback_(ePtr);
TransactionRollback("The transaction has been rolled back"));
for (auto const &cmd : sqlCmdBuffer_)
{
if (cmd->exceptionCallback_)
{
cmd->exceptionCallback_(exceptPtr);
}
}
sqlCmdBuffer_.clear();
}
if (usedUpCallback_)
{
usedUpCallback_();
usedUpCallback_ = std::function<void()>();
} }
} }
} }
const char *TransactionImpl::beginSql() const noexcept
{
if (type_ != ClientType::Sqlite3)
return "begin";
switch (transactionType_)
{
case TransactionType::Immediate:
return "begin immediate";
case TransactionType::Exclusive:
return "begin exclusive";
default:
return "begin";
}
}
void TransactionImpl::doBegin() void TransactionImpl::doBegin()
{ {
loop_->queueInLoop([thisPtr = shared_from_this()]() { loop_->queueInLoop([thisPtr = shared_from_this()]() {
@@ -280,7 +308,7 @@ void TransactionImpl::doBegin()
thisPtr->isWorking_ = true; thisPtr->isWorking_ = true;
thisPtr->thisPtr_ = thisPtr; thisPtr->thisPtr_ = thisPtr;
thisPtr->connectionPtr_->execSql( thisPtr->connectionPtr_->execSql(
"begin", thisPtr->beginSql(),
0, 0,
{}, {},
{}, {},
@@ -289,6 +317,12 @@ void TransactionImpl::doBegin()
[thisPtr](const std::exception_ptr &) { [thisPtr](const std::exception_ptr &) {
LOG_ERROR << "Error occurred in transaction begin"; LOG_ERROR << "Error occurred in transaction begin";
thisPtr->isCommitedOrRolledback_ = true; thisPtr->isCommitedOrRolledback_ = true;
thisPtr->isWorking_ = false;
thisPtr->thisPtr_.reset();
thisPtr->failBufferedCommands(std::make_exception_ptr(
TransactionRollback("Transaction begin failed, cannot "
"execute queued SQL")));
thisPtr->releaseConnection();
}); });
}); });
} }
+10 -3
View File
@@ -30,7 +30,8 @@ class TransactionImpl : public Transaction,
TransactionImpl(ClientType type, TransactionImpl(ClientType type,
const DbConnectionPtr &connPtr, const DbConnectionPtr &connPtr,
std::function<void(bool)> commitCallback, std::function<void(bool)> commitCallback,
std::function<void()> usedUpCallback); std::function<void()> usedUpCallback,
TransactionType transType = TransactionType::Deferred);
~TransactionImpl() override; ~TransactionImpl() override;
void rollback() override; void rollback() override;
@@ -113,14 +114,16 @@ class TransactionImpl : public Transaction,
std::function<void(const std::exception_ptr &)> &&exceptCallback); std::function<void(const std::exception_ptr &)> &&exceptCallback);
std::shared_ptr<Transaction> newTransaction( std::shared_ptr<Transaction> newTransaction(
const std::function<void(bool)> &) noexcept(false) override const std::function<void(bool)> &,
TransactionType) noexcept(false) override
{ {
return shared_from_this(); return shared_from_this();
} }
void newTransactionAsync( void newTransactionAsync(
const std::function<void(const std::shared_ptr<Transaction> &)> const std::function<void(const std::shared_ptr<Transaction> &)>
&callback) override &callback,
TransactionType) override
{ {
callback(shared_from_this()); callback(shared_from_this());
} }
@@ -129,6 +132,8 @@ class TransactionImpl : public Transaction,
bool isCommitedOrRolledback_{false}; bool isCommitedOrRolledback_{false};
bool isWorking_{false}; bool isWorking_{false};
void execNewTask(); void execNewTask();
void releaseConnection();
void failBufferedCommands(const std::exception_ptr &ePtr);
struct SqlCmd struct SqlCmd
{ {
@@ -149,10 +154,12 @@ class TransactionImpl : public Transaction,
friend class DbClientImpl; friend class DbClientImpl;
friend class DbClientLockFree; friend class DbClientLockFree;
void doBegin(); void doBegin();
const char *beginSql() const noexcept;
trantor::EventLoop *loop_; trantor::EventLoop *loop_;
std::function<void(bool)> commitCallback_; std::function<void(bool)> commitCallback_;
std::shared_ptr<TransactionImpl> thisPtr_; std::shared_ptr<TransactionImpl> thisPtr_;
double timeout_{-1.0}; double timeout_{-1.0};
TransactionType transactionType_{TransactionType::Deferred};
}; };
} // namespace orm } // namespace orm
} // namespace drogon } // namespace drogon
+185 -6
View File
@@ -2741,11 +2741,9 @@ DROGON_TEST(MySQLTest)
#endif #endif
#if USE_SQLITE3 #if USE_SQLITE3
DbClientPtr sqlite3Client;
DROGON_TEST(SQLite3Test) DROGON_TEST(SQLite3Test)
{ {
auto &clientPtr = sqlite3Client; auto clientPtr = DbClient::newSqlite3Client("filename=:memory:", 1);
REQUIRE(clientPtr != nullptr); REQUIRE(clientPtr != nullptr);
// Prepare the test environment // Prepare the test environment
@@ -4063,6 +4061,190 @@ DROGON_TEST(SQLite3Test)
} }
#endif #endif
#if USE_SQLITE3
DROGON_TEST(SQLite3TransactionTypeTest)
{
auto clientPtr = DbClient::newSqlite3Client("filename=:memory:", 1);
REQUIRE(clientPtr != nullptr);
// Ensure the test table exists
try
{
clientPtr->execSqlSync(
"CREATE TABLE IF NOT EXISTS trans_type_test "
"(id INTEGER PRIMARY KEY, val INTEGER NOT NULL)");
clientPtr->execSqlSync("DELETE FROM trans_type_test");
}
catch (const DrogonDbException &e)
{
FAULT("sqlite3 - TransactionType setup what():", e.base().what());
return;
}
// --- Deferred (default) ---
{
try
{
auto trans = clientPtr->newTransaction(TransactionType::Deferred);
trans->execSqlSync(
"INSERT INTO trans_type_test(id, val) VALUES(1, 10)");
// trans commits on destruction
}
catch (const DrogonDbException &e)
{
FAULT("sqlite3 - TransactionType::Deferred what():",
e.base().what());
return;
}
auto r = clientPtr->execSqlSync(
"SELECT val FROM trans_type_test WHERE id=1");
MANDATE(r.size() == 1);
MANDATE(r[0][0].as<int>() == 10);
SUCCESS();
}
// --- Immediate ---
{
try
{
auto trans = clientPtr->newTransaction(TransactionType::Immediate);
trans->execSqlSync(
"INSERT INTO trans_type_test(id, val) VALUES(2, 20)");
}
catch (const DrogonDbException &e)
{
FAULT("sqlite3 - TransactionType::Immediate what():",
e.base().what());
return;
}
auto r = clientPtr->execSqlSync(
"SELECT val FROM trans_type_test WHERE id=2");
MANDATE(r.size() == 1);
MANDATE(r[0][0].as<int>() == 20);
SUCCESS();
}
// --- Exclusive ---
{
try
{
auto trans = clientPtr->newTransaction(TransactionType::Exclusive);
trans->execSqlSync(
"INSERT INTO trans_type_test(id, val) VALUES(3, 30)");
}
catch (const DrogonDbException &e)
{
FAULT("sqlite3 - TransactionType::Exclusive what():",
e.base().what());
return;
}
auto r = clientPtr->execSqlSync(
"SELECT val FROM trans_type_test WHERE id=3");
MANDATE(r.size() == 1);
MANDATE(r[0][0].as<int>() == 30);
SUCCESS();
}
// --- Rollback works correctly with Immediate ---
{
try
{
auto trans = clientPtr->newTransaction(TransactionType::Immediate);
trans->execSqlSync(
"INSERT INTO trans_type_test(id, val) VALUES(99, 99)");
trans->rollback();
}
catch (const DrogonDbException &e)
{
FAULT("sqlite3 - TransactionType::Immediate rollback what():",
e.base().what());
return;
}
auto r = clientPtr->execSqlSync(
"SELECT val FROM trans_type_test WHERE id=99");
MANDATE(r.size() == 0);
SUCCESS();
}
}
// Verify the locking mode is actually used by testing observable SQLite
// locking behaviour. BEGIN IMMEDIATE acquires a RESERVED lock upfront, so a
// second concurrent BEGIN IMMEDIATE on another connection to the same
// database must fail with SQLITE_BUSY. If plain BEGIN were used instead, the
// second connection would succeed (only a SHARED lock is held until the first
// write).
DROGON_TEST(SQLite3TransactionTypeLockingTest)
{
// A pool of 2 connections to a shared file-based database gives us two
// independent SQLite connections that observe each other's locks.
const auto nonce =
std::chrono::steady_clock::now().time_since_epoch().count();
const auto dbPath =
"drogon_trans_type_lock_test_" + std::to_string(nonce) + ".db";
std::remove(dbPath.c_str());
auto pool = DbClient::newSqlite3Client("filename=" + dbPath, 2);
// WAL mode is required: it changes BEGIN IMMEDIATE from acquiring a
// RESERVED lock to acquiring the WAL write lock. This matches production
// usage and makes the busy semantics more predictable — only one writer
// is ever permitted and SQLITE_BUSY is returned immediately (no timeout
// retry) when a second BEGIN IMMEDIATE is attempted.
pool->execSqlSync("PRAGMA journal_mode=WAL");
// No retry delay: SQLITE_BUSY must surface as an exception immediately.
pool->execSqlSync("PRAGMA busy_timeout=0");
pool->execSqlSync(
"CREATE TABLE IF NOT EXISTS lock_test (id INTEGER PRIMARY KEY)");
std::shared_ptr<Transaction> transA;
// Hold an IMMEDIATE transaction on connection A.
try
{
transA = pool->newTransaction(TransactionType::Immediate);
// doBegin() is asynchronous — the BEGIN IMMEDIATE is queued to the
// connection's event loop. Run a synchronous query through the
// transaction to flush the queue; once execSqlSync returns, the
// RESERVED lock is definitely held.
transA->execSqlSync("SELECT 1");
}
catch (const DrogonDbException &e)
{
std::remove(dbPath.c_str());
FAULT("sqlite3 - TransactionType::Immediate locking setup what():",
e.base().what());
return;
}
// Connection B attempting BEGIN IMMEDIATE must fail because A already
// holds the RESERVED lock. SQLite's default busy_timeout is 0.
bool gotBusy = false;
try
{
auto transB = pool->newTransaction(TransactionType::Immediate);
transB->execSqlSync("SELECT 1");
transB->rollback();
}
catch (const DrogonDbException &)
{
gotBusy = true;
}
transA->rollback();
std::remove(dbPath.c_str());
if (gotBusy)
{
SUCCESS();
}
else
{
FAULT(
"sqlite3 - TransactionType::Immediate locking: second BEGIN "
"IMMEDIATE should have failed while the first was held, but it "
"succeeded. This means BEGIN IMMEDIATE is not being sent.");
}
}
#endif
using namespace drogon; using namespace drogon;
int main(int argc, char **argv) int main(int argc, char **argv)
@@ -4079,9 +4261,6 @@ int main(int argc, char **argv)
"client_encoding=utf8", "client_encoding=utf8",
1, 1,
true); true);
#endif
#if USE_SQLITE3
sqlite3Client = DbClient::newSqlite3Client("filename=:memory:", 1);
#endif #endif
const int testStatus = test::run(argc, argv); const int testStatus = test::run(argc, argv);
return testStatus; return testStatus;
+134 -32
View File
@@ -23,46 +23,147 @@ else
fi fi
echo "drogon_ctl_exec: " ${drogon_ctl_exec} echo "drogon_ctl_exec: " ${drogon_ctl_exec}
#Make integration_test_server run as a daemon if [ "X$os" = "Xwindows" ]; then
integration_test_client_exec=./integration_test_client.exe
integration_test_server_exec=./integration_test_server.exe
else
integration_test_client_exec=./integration_test_client
integration_test_server_exec=./integration_test_server
fi
function update_config_line()
{
local key="$1"
local value="$2"
local file="$3"
sed -i.bak -e "s/\"${key}\".*$/\"${key}\": ${value},/" "$file"
rm -f "$file.bak"
}
function cleanup_integration_test_server()
{
if [ "X$os" = "Xwindows" ]; then
taskkill //F //IM integration_test_server.exe > /dev/null 2>&1 || true
else
killall integration_test_server > /dev/null 2>&1 || true
pkill -f '/integration_test_server$' > /dev/null 2>&1 || true
fi
}
function wait_for_url()
{
local url="$1"
local timeout_seconds="$2"
local curl_args=(--silent --show-error --output /dev/null --max-time 2)
if [ "$url" != "${url#https://}" ]; then
curl_args+=(--insecure)
fi
local attempt=0
while [ $attempt -lt $timeout_seconds ]; do
if curl "${curl_args[@]}" "$url"; then
return 0
fi
attempt=$((attempt + 1))
sleep 1
done
return 1
}
function wait_for_integration_test_server()
{
local server_pid="$1"
local server_log="$2"
if ! wait_for_url "http://127.0.0.1:8848/" 30; then
echo "Timed out waiting for integration_test_server to accept HTTP requests"
if kill -0 "$server_pid" > /dev/null 2>&1; then
echo "integration_test_server is still running, recent log output:"
else
echo "integration_test_server exited before becoming ready, recent log output:"
fi
if [ -f "$server_log" ]; then
tail -n 50 "$server_log"
fi
return 1
fi
wait_for_url "https://127.0.0.1:8849/" 5 > /dev/null 2>&1 || true
return 0
}
function get_cpu_count()
{
if command -v nproc > /dev/null 2>&1; then
nproc
return
fi
if command -v getconf > /dev/null 2>&1; then
getconf _NPROCESSORS_ONLN
return
fi
if command -v sysctl > /dev/null 2>&1; then
sysctl -n hw.logicalcpu
return
fi
echo 1
}
trap cleanup_integration_test_server EXIT
# Run the integration test server in the background and wait until it is ready.
function do_integration_test() function do_integration_test()
{ {
pushd $test_root pushd "$test_root"
if [ "X$os" = "Xlinux" ]; then update_config_line "run_as_daemon" "false" config.example.json
sed -i -e "s/\"run_as_daemon.*$/\"run_as_daemon\": true\,/" config.example.json update_config_line "relaunch_on_error" "false" config.example.json
fi update_config_line "number_of_threads" "1" config.example.json
sed -i -e "s/\"relaunch_on_error.*$/\"relaunch_on_error\": true\,/" config.example.json update_config_line "use_brotli" "true" config.example.json
sed -i -e "s/\"threads_num.*$/\"threads_num\": 0\,/" config.example.json
sed -i -e "s/\"use_brotli.*$/\"use_brotli\": true\,/" config.example.json
if [ "$1" = "stream_mode" ]; then if [ "$1" = "stream_mode" ]; then
sed -i -e "s/\"enable_request_stream.*$/\"enable_request_stream\": true\,/" config.example.json update_config_line "enable_request_stream" "true" config.example.json
else else
sed -i -e "s/\"enable_request_stream.*$/\"enable_request_stream\": false\,/" config.example.json update_config_line "enable_request_stream" "false" config.example.json
fi fi
if [ ! -f "integration_test_client" ]; then if [ ! -f "$integration_test_client_exec" ]; then
echo "Build failed" echo "Build failed"
exit -1 exit -1
fi fi
if [ ! -f "integration_test_server" ]; then if [ ! -f "$integration_test_server_exec" ]; then
echo "Build failed" echo "Build failed"
exit -1 exit -1
fi fi
killall -9 integration_test_server cleanup_integration_test_server
./integration_test_server &
sleep 4 local server_log=integration_test_server.log
rm -f "$server_log"
"$integration_test_server_exec" > "$server_log" 2>&1 &
local server_pid=$!
if ! wait_for_integration_test_server "$server_pid" "$server_log"; then
exit -1
fi
echo "Running the integration test $1" echo "Running the integration test $1"
./integration_test_client -s "$integration_test_client_exec" -s
if [ $? -ne 0 ]; then if [ $? -ne 0 ]; then
echo "Integration test failed $1" echo "Integration test failed $1"
if [ -f "$server_log" ]; then
tail -n 50 "$server_log"
fi
exit -1 exit -1
fi fi
killall -9 integration_test_server cleanup_integration_test_server
popd popd
} }
@@ -70,7 +171,7 @@ function do_integration_test()
function do_drogon_ctl_test() function do_drogon_ctl_test()
{ {
echo "Testing drogon_ctl" echo "Testing drogon_ctl"
pushd $test_root pushd "$test_root"
rm -rf drogon_test rm -rf drogon_test
${drogon_ctl_exec} create project drogon_test ${drogon_ctl_exec} create project drogon_test
@@ -122,22 +223,23 @@ function do_drogon_ctl_test()
make_flags='' make_flags=''
cmake_gen='' cmake_gen=''
parallel=1 parallel=1
cpu_count=$(get_cpu_count)
# simulate ninja's parallelism # simulate ninja's parallelism
case $(nproc) in case $cpu_count in
1) 1)
parallel=$(($(nproc) + 1)) parallel=$((cpu_count + 1))
;; ;;
2) 2)
parallel=$(($(nproc) + 1)) parallel=$((cpu_count + 1))
;; ;;
*) *)
parallel=$(($(nproc) + 2)) parallel=$((cpu_count + 2))
;; ;;
esac esac
if [ "X$os" = "Xlinux" ]; then if [ "X$os" = "Xlinux" ]; then
if [ -f /bin/ninja ]; then if command -v ninja > /dev/null 2>&1; then
cmake_gen='-G Ninja' cmake_gen='-G Ninja'
else else
make_flags="$make_flags -j$parallel" make_flags="$make_flags -j$parallel"
@@ -162,11 +264,11 @@ function do_drogon_ctl_test()
exit -1 exit -1
fi fi
if [ "X$os" = "Xlinux" ]; then if [ "X$os" = "Xlinux" ]; then
if [ ! -f "drogon_test" ]; then if [ ! -f "drogon_test" ]; then
echo "Failed to build drogon_test" echo "Failed to build drogon_test"
exit -1 exit -1
fi fi
else else
if [ ! -f "Debug\drogon_test.exe" ]; then if [ ! -f "Debug\drogon_test.exe" ]; then
echo "Failed to build drogon_test" echo "Failed to build drogon_test"
@@ -183,7 +285,7 @@ function do_drogon_ctl_test()
function do_unittest() function do_unittest()
{ {
echo "Unit testing" echo "Unit testing"
pushd $src_dir/build pushd "$src_dir/build"
ctest . --output-on-failure ctest . --output-on-failure
if [ $? -ne 0 ]; then if [ $? -ne 0 ]; then
@@ -195,7 +297,7 @@ function do_unittest()
function do_db_test() function do_db_test()
{ {
pushd $src_dir/build pushd "$src_dir/build"
if [ -f "./orm_lib/tests/db_test" ]; then if [ -f "./orm_lib/tests/db_test" ]; then
echo "Test database" echo "Test database"
./orm_lib/tests/db_test -s ./orm_lib/tests/db_test -s
@@ -246,9 +348,9 @@ function do_db_test()
fi fi
} }
if ! drogon_ctl -v > /dev/null 2>&1 if [ ! -f "$drogon_ctl_exec" ]
then then
echo "Warning: No drogon_ctl, skip integration test and drogon_ctl test" echo "Warning: No built drogon_ctl, skip integration test and drogon_ctl test"
else else
do_integration_test do_integration_test
do_integration_test stream_mode do_integration_test stream_mode
Vendored Executable → Regular
+12 -7
View File
@@ -5,8 +5,17 @@ option(BUILD_DOC "Build Doxygen documentation" OFF)
option(BUILD_C-ARES "Build C-ARES" ON) option(BUILD_C-ARES "Build C-ARES" ON)
option(BUILD_TESTING "Build tests" OFF) option(BUILD_TESTING "Build tests" OFF)
option(BUILD_SHARED_LIBS "Build trantor as a shared lib" OFF) option(BUILD_SHARED_LIBS "Build trantor as a shared lib" OFF)
option(TRANTOR_USE_TLS set(TRANTOR_USE_TLS
"TLS provider for trantor. Valid options are 'openssl', 'botan' or '' (let the build scripr decide)" "" ""
CACHE STRING "TLS provider for trantor. Valid options are 'openssl', 'botan', 'none' or '' (auto-detect)"
)
set_property(
CACHE TRANTOR_USE_TLS
PROPERTY STRINGS
""
openssl
botan
none
) )
option(USE_SPDLOG "Allow using the spdlog logging library" OFF) option(USE_SPDLOG "Allow using the spdlog logging library" OFF)
@@ -14,7 +23,7 @@ list(APPEND CMAKE_MODULE_PATH ${PROJECT_SOURCE_DIR}/cmake_modules/)
set(TRANTOR_MAJOR_VERSION 1) set(TRANTOR_MAJOR_VERSION 1)
set(TRANTOR_MINOR_VERSION 5) set(TRANTOR_MINOR_VERSION 5)
set(TRANTOR_PATCH_VERSION 26) set(TRANTOR_PATCH_VERSION 28)
set(TRANTOR_VERSION ${TRANTOR_MAJOR_VERSION}.${TRANTOR_MINOR_VERSION}.${TRANTOR_PATCH_VERSION}) set(TRANTOR_VERSION ${TRANTOR_MAJOR_VERSION}.${TRANTOR_MINOR_VERSION}.${TRANTOR_PATCH_VERSION})
include(GNUInstallDirs) include(GNUInstallDirs)
@@ -158,10 +167,6 @@ else(WIN32)
set(TRANTOR_SOURCES ${TRANTOR_SOURCES} trantor/net/inner/FileBufferNodeUnix.cc) set(TRANTOR_SOURCES ${TRANTOR_SOURCES} trantor/net/inner/FileBufferNodeUnix.cc)
endif(WIN32) endif(WIN32)
# Somehow the default value of TRANTOR_USE_TLS is OFF
if(TRANTOR_USE_TLS STREQUAL OFF)
set(TRANTOR_USE_TLS "")
endif()
set(VALID_TLS_PROVIDERS "openssl" "botan" "none") set(VALID_TLS_PROVIDERS "openssl" "botan" "none")
list( list(
FIND FIND
+29 -1
View File
@@ -4,6 +4,30 @@ All notable changes to this project will be documented in this file.
## [Unreleased] ## [Unreleased]
## [1.5.28] - 2026-05-06
### Fixed
- Avoid abort on closeWrite shutdown failure.
## [1.5.27] - 2026-05-06
### Changed
- Add automatic SSL.
- Add getter for `TcpConnection::closeCallback_`.
- Remove spurious executable permissions from non-script sources.
### Fixed
- Fix `TRANTOR_USE_TLS` cache setting in CMake.
- Fix TLS implementation quirks.
- Fix server-side mTLS client certificate hostname validation.
## [1.5.26] - 2026-01-26 ## [1.5.26] - 2026-01-26
### Changed ### Changed
@@ -742,7 +766,11 @@ All notable changes to this project will be documented in this file.
## [1.0.0-rc1] - 2019-06-11 ## [1.0.0-rc1] - 2019-06-11
[Unreleased]: https://github.com/an-tao/trantor/compare/v1.5.26...HEAD [Unreleased]: https://github.com/an-tao/trantor/compare/v1.5.28...HEAD
[1.5.28]: https://github.com/an-tao/trantor/compare/v1.5.27...v1.5.28
[1.5.27]: https://github.com/an-tao/trantor/compare/v1.5.26...v1.5.27
[1.5.26]: https://github.com/an-tao/trantor/compare/v1.5.25...v1.5.26 [1.5.26]: https://github.com/an-tao/trantor/compare/v1.5.25...v1.5.26
Vendored Executable → Regular
View File
@@ -351,6 +351,10 @@ class TRANTOR_EXPORT TcpConnection
{ {
closeCallback_ = std::move(cb); closeCallback_ = std::move(cb);
} }
CloseCallback getCloseCallback() const
{
return closeCallback_;
}
void setSSLErrorCallback(const SSLErrorCallback &cb) void setSSLErrorCallback(const SSLErrorCallback &cb)
{ {
sslErrorCallback_ = cb; sslErrorCallback_ = cb;
@@ -367,6 +371,8 @@ class TRANTOR_EXPORT TcpConnection
size_t timeout, size_t timeout,
const std::shared_ptr<TimingWheel> &timingWheel) = 0; const std::shared_ptr<TimingWheel> &timingWheel) = 0;
virtual void forwardToTLSBuffer(MsgBuffer *buffer) = 0;
protected: protected:
// callbacks // callbacks
RecvMessageCallback recvMsgCallback_; RecvMessageCallback recvMsgCallback_;
View File
@@ -203,6 +203,12 @@ class TcpConnectionImpl : public TcpConnection,
timingWheel->insertEntry(timeout, entry); timingWheel->insertEntry(timeout, entry);
} }
void forwardToTLSBuffer(MsgBuffer *buffer) override
{
if (tlsProviderPtr_)
tlsProviderPtr_->recvData(buffer);
}
private: private:
/// Internal use only. /// Internal use only.
@@ -474,8 +474,6 @@ SSLContextPtr trantor::newSSLContext(const TLSPolicy &policy, bool server)
ctx->certStore = ctx->certStore =
std::make_shared<Botan::Flatfile_Certificate_Store>( std::make_shared<Botan::Flatfile_Certificate_Store>(
policy.getCaPath()); policy.getCaPath());
if (server)
ctx->requireClientCert = true;
} }
else if (policy.getUseSystemCertStore()) else if (policy.getUseSystemCertStore())
{ {
@@ -484,6 +482,8 @@ SSLContextPtr trantor::newSSLContext(const TLSPolicy &policy, bool server)
ctx->certStore = systemCertStore; ctx->certStore = systemCertStore;
} }
} }
if (server && policy.getValidate() && !policy.getCaPath().empty())
ctx->requireClientCert = true;
if (policy.getUseOldTLS()) if (policy.getUseOldTLS())
LOG_WARN << "SSLPloicy have set useOldTLS to true. BUt Botan does not " LOG_WARN << "SSLPloicy have set useOldTLS to true. BUt Botan does not "
@@ -8,12 +8,10 @@
#include <openssl/bio.h> #include <openssl/bio.h>
#include <openssl/x509v3.h> #include <openssl/x509v3.h>
#include <fstream>
#include <memory> #include <memory>
#include <mutex> #include <mutex>
#include <list> #include <list>
#include <unordered_map> #include <unordered_map>
#include <array>
#include <limits> #include <limits>
#include "callbacks.h" #include "callbacks.h"
@@ -70,62 +68,6 @@ inline bool loadWindowsSystemCert(X509_STORE *store)
} }
#endif #endif
inline bool verifyCommonName(X509 *cert, const std::string &hostname)
{
X509_NAME *subjectName = X509_get_subject_name(cert);
if (subjectName != nullptr)
{
std::array<char, BUFSIZ> name;
auto length = X509_NAME_get_text_by_NID(subjectName,
NID_commonName,
name.data(),
(int)name.size());
if (length == -1)
return false;
return utils::verifySslName(std::string(name.begin(),
name.begin() + length),
hostname);
}
return false;
}
inline bool verifyAltName(X509 *cert, const std::string &hostname)
{
bool good = false;
auto altNames = static_cast<const struct stack_st_GENERAL_NAME *>(
X509_get_ext_d2i(cert, NID_subject_alt_name, nullptr, nullptr));
if (altNames)
{
int numNames = sk_GENERAL_NAME_num(altNames);
for (int i = 0; i < numNames && !good; i++)
{
auto val = sk_GENERAL_NAME_value(altNames, i);
if (val->type != GEN_DNS)
{
LOG_WARN << "Name using IP addresses are not supported. Open "
"an issue if you need that feature";
continue;
}
#if (OPENSSL_VERSION_NUMBER >= 0x10100000L)
auto name = (const char *)ASN1_STRING_get0_data(val->d.ia5);
#else
auto name = (const char *)ASN1_STRING_data(val->d.ia5);
#endif
auto name_len = (size_t)ASN1_STRING_length(val->d.ia5);
good = utils::verifySslName(std::string(name, name + name_len),
hostname);
}
}
GENERAL_NAMES_free((STACK_OF(GENERAL_NAME) *)altNames);
return good;
}
static bool validatePeerCertificate(SSL *ssl, static bool validatePeerCertificate(SSL *ssl,
X509 *cert, X509 *cert,
const std::string &hostname, const std::string &hostname,
@@ -136,12 +78,16 @@ static bool validatePeerCertificate(SSL *ssl,
assert(cert != nullptr); assert(cert != nullptr);
LOG_TRACE << "Validating peer certificate"; LOG_TRACE << "Validating peer certificate";
if (isServer) if (!isServer)
{ {
bool domainIsValid = const int rc =
verifyCommonName(cert, hostname) || verifyAltName(cert, hostname); X509_check_host(cert, hostname.data(), hostname.size(), 0, nullptr);
if (!domainIsValid) if (rc != 1)
{
LOG_TRACE << "Peer certificate does not match hostname: "
<< hostname;
return false; return false;
}
} }
auto result = SSL_get_verify_result(ssl); auto result = SSL_get_verify_result(ssl);
@@ -423,16 +369,20 @@ class SessionManager
#endif #endif
} }
// Returns a session with an additional reference held by the caller.
// Caller must SSL_SESSION_free() when done. Required because the entry
// in sessionMap_ may be evicted/replaced/expired by another thread the
// moment we release the mutex, so the SessionManager's reference is not
// a stable ownership root for the returned pointer.
SSL_SESSION *get(const std::string &hostname, InetAddress peerAddr) SSL_SESSION *get(const std::string &hostname, InetAddress peerAddr)
{ {
std::lock_guard<std::mutex> lock(mutex_); std::lock_guard<std::mutex> lock(mutex_);
auto key = toKey(hostname, peerAddr); auto it = sessionMap_.find(toKey(hostname, peerAddr));
auto it = sessionMap_.find(key); if (it == sessionMap_.end())
if (it != sessionMap_.end()) return nullptr;
{ SSL_SESSION *s = it->second->session;
return it->second->session; SSL_SESSION_up_ref(s);
} return s;
return nullptr;
} }
void removeExcessSession() void removeExcessSession()
@@ -529,7 +479,9 @@ struct OpenSSLProvider : public TLSProvider, public NonCopyable
conn_->peerAddr()); conn_->peerAddr());
if (cachedSession) if (cachedSession)
{ {
// SSL_set_session takes its own reference; release ours.
SSL_set_session(ssl_, cachedSession); SSL_set_session(ssl_, cachedSession);
SSL_SESSION_free(cachedSession);
} }
SSL_set_connect_state(ssl_); SSL_set_connect_state(ssl_);
} }
@@ -671,7 +623,10 @@ struct OpenSSLProvider : public TLSProvider, public NonCopyable
cert, cert,
policyPtr_->getHostname(), policyPtr_->getHostname(),
policyPtr_->getAllowBrokenChain(), policyPtr_->getAllowBrokenChain(),
contextPtr_->isServer); !contextPtr_
->isServer); // From the server's point of view,
// the client certificate is verified
// and vice versa
if (!valid) if (!valid)
{ {
LOG_TRACE LOG_TRACE
@@ -0,0 +1,54 @@
#include <trantor/net/TcpClient.h>
#include <trantor/utils/Logger.h>
#include <trantor/net/EventLoopThread.h>
#include <string>
#include <iostream>
#include <atomic>
using namespace trantor;
#define USE_IPV6 0
int main()
{
trantor::Logger::setLogLevel(trantor::Logger::kDebug);
LOG_DEBUG << "TcpClient class test!";
EventLoop loop;
#if USE_IPV6
InetAddress serverAddr("::1", 8888, true);
#else
InetAddress serverAddr("127.0.0.1", 8888);
#endif
std::shared_ptr<trantor::TcpClient> client[10];
std::atomic_int connCount;
connCount = 1;
for (int i = 0; i < 1; ++i)
{
client[i] = std::make_shared<trantor::TcpClient>(&loop,
serverAddr,
"tcpclienttest");
auto policy = TLSPolicy::defaultClientPolicy();
policy->setValidate(false);
client[i]->enableSSL(std::move(policy));
client[i]->setConnectionCallback(
[i, &loop, &connCount](const TcpConnectionPtr &conn) {
if (conn->connected())
{
LOG_DEBUG << i << " connected";
conn->send("Hello");
}
else
{
LOG_DEBUG << i << " disconnected";
--connCount;
if (connCount == 0)
loop.quit();
}
});
client[i]->setMessageCallback(
[](const TcpConnectionPtr &conn, MsgBuffer *buf) {
auto msg = std::string(buf->peek(), buf->readableBytes());
LOG_INFO << msg;
buf->retrieveAll();
});
client[i]->connect();
}
loop.loop();
}
@@ -0,0 +1,63 @@
#include <trantor/net/TcpServer.h>
#include <trantor/utils/Logger.h>
#include <trantor/net/EventLoopThread.h>
#include <string>
#include <iostream>
using namespace trantor;
#define USE_IPV6 0
bool has_ssl(MsgBuffer *buffer)
{
if (buffer->readableBytes() < 3)
return false;
const char *data = buffer->peek();
unsigned char byte1 = static_cast<unsigned char>(data[0]);
unsigned char byte2 = static_cast<unsigned char>(data[1]);
unsigned char byte3 = static_cast<unsigned char>(data[2]);
return (byte1 == 0x16) && (byte2 == 0x03) && (byte3 == 0x01);
}
int main()
{
LOG_DEBUG << "test start";
Logger::setLogLevel(Logger::kDebug);
EventLoopThread loopThread;
loopThread.run();
#if USE_IPV6
InetAddress addr(8888, true, true);
#else
InetAddress addr(8888);
#endif
TcpServer server(loopThread.getLoop(), addr, "test");
// auto ctx = newSSLServerContext("server.pem", "server.pem", {});
LOG_INFO << "start";
server.setRecvMessageCallback(
[](const TcpConnectionPtr &connectionPtr, MsgBuffer *buffer) {
if (has_ssl(buffer))
{
LOG_DEBUG << "SSL data received";
auto policy =
TLSPolicy::defaultServerPolicy("server.crt", "server.key");
connectionPtr->startEncryption(policy, true);
connectionPtr->forwardToTLSBuffer(buffer);
return;
}
LOG_DEBUG << std::string{buffer->peek(), buffer->readableBytes()};
connectionPtr->send(*buffer);
buffer->retrieveAll();
connectionPtr->shutdown();
});
server.setConnectionCallback([](const TcpConnectionPtr &connPtr) {
if (connPtr->connected())
{
LOG_DEBUG << "New connection";
}
else if (connPtr->disconnected())
{
LOG_DEBUG << "connection disconnected";
}
});
server.setIoLoopNum(3);
server.start();
loopThread.wait();
}
@@ -23,6 +23,8 @@ add_executable(logger_macro_test LoggerMacroTest.cc)
add_executable(delayed_ssl_server_test DelayedSSLServerTest.cc) add_executable(delayed_ssl_server_test DelayedSSLServerTest.cc)
add_executable(delayed_ssl_client_test DelayedSSLClientTest.cc) add_executable(delayed_ssl_client_test DelayedSSLClientTest.cc)
add_executable(tcp_asyncstream_server_test TcpAsyncStreamServerTest.cc) add_executable(tcp_asyncstream_server_test TcpAsyncStreamServerTest.cc)
add_executable(automatic_ssl_server_test AutomaticSSLServerTest.cc)
add_executable(automatic_ssl_client_test AutomaticSSLClientTest.cc)
set(targets_list set(targets_list
ssl_server_test ssl_server_test
ssl_client_test ssl_client_test
@@ -49,6 +51,8 @@ set(targets_list
delayed_ssl_server_test delayed_ssl_server_test
delayed_ssl_client_test delayed_ssl_client_test
tcp_asyncstream_server_test tcp_asyncstream_server_test
automatic_ssl_server_test
automatic_ssl_client_test
) )
if(HAVE_SPDLOG) if(HAVE_SPDLOG)
View File
@@ -4,7 +4,6 @@ add_executable(inetaddress_unittest InetAddressUnittest.cc)
add_executable(date_unittest DateUnittest.cc) add_executable(date_unittest DateUnittest.cc)
add_executable(split_string_unittest splitStringUnittest.cc) add_executable(split_string_unittest splitStringUnittest.cc)
add_executable(string_encoding_unittest stringEncodingUnittest.cc) add_executable(string_encoding_unittest stringEncodingUnittest.cc)
add_executable(ssl_name_verify_unittest sslNameVerifyUnittest.cc)
add_executable(hash_unittest HashUnittest.cc) add_executable(hash_unittest HashUnittest.cc)
set(UNITTEST_TARGETS set(UNITTEST_TARGETS
msgbuffer_unittest msgbuffer_unittest
@@ -12,7 +11,6 @@ set(UNITTEST_TARGETS
date_unittest date_unittest
split_string_unittest split_string_unittest
string_encoding_unittest string_encoding_unittest
ssl_name_verify_unittest
hash_unittest hash_unittest
) )
set_property(TARGET ${UNITTEST_TARGETS} PROPERTY CXX_STANDARD 14) set_property(TARGET ${UNITTEST_TARGETS} PROPERTY CXX_STANDARD 14)
@@ -1,50 +0,0 @@
#include <trantor/utils/Utilities.h>
#include <gtest/gtest.h>
#include <iostream>
using namespace trantor;
using namespace trantor::utils;
TEST(sslNameCheck, baseCases)
{
EXPECT_EQ(verifySslName("example.com", "example.com"), true);
EXPECT_EQ(verifySslName("example.com", "example.org"), false);
EXPECT_EQ(verifySslName("example.com", "www.example.com"), false);
}
TEST(sslNameCheck, rfc6125Examples)
{
EXPECT_EQ(verifySslName("*.example.com", "foo.example.com"), true);
EXPECT_EQ(verifySslName("*.example.com", "foo.bar.example.com"), false);
EXPECT_EQ(verifySslName("*.example.com", "example.com"), false);
EXPECT_EQ(verifySslName("*bar.example.com", "foobar.example.com"), true);
EXPECT_EQ(verifySslName("baz*.example.com", "baz1.example.com"), true);
EXPECT_EQ(verifySslName("b*z.example.com", "buzz.example.com"), true);
}
TEST(sslNameCheck, rfcCounterExamples)
{
EXPECT_EQ(verifySslName("buz*.example.com", "buaz.example.com"), false);
EXPECT_EQ(verifySslName("*bar.example.com", "aaasdasbaz.example.com"),
false);
EXPECT_EQ(verifySslName("b*z.example.com", "baaaaaa.example.com"), false);
}
TEST(sslNameCheck, wildExamples)
{
EXPECT_EQ(verifySslName("datatracker.ietf.org", "datatracker.ietf.org"),
true);
EXPECT_EQ(verifySslName("*.nsysu.edu.tw", "nsysu.edu.tw"), false);
EXPECT_EQ(verifySslName("nsysu.edu.tw", "nsysu.edu.tw"), true);
}
TEST(sslNameCheck, edgeCase)
{
EXPECT_EQ(verifySslName(".example.com", "example.com"), false);
EXPECT_EQ(verifySslName("example.com.", "example.com."), true);
}
int main(int argc, char **argv)
{
testing::InitGoogleTest(&argc, argv);
return RUN_ALL_TESTS();
}
View File
@@ -228,120 +228,6 @@ std::string fromWidePath(const std::wstring &wstrPath)
return toUtf8(srcPath); return toUtf8(srcPath);
} }
bool verifySslName(const std::string &certName, const std::string &hostname)
{
if (certName.find('*') == std::string::npos)
{
return certName == hostname;
}
size_t firstDot = certName.find('.');
size_t hostFirstDot = hostname.find('.');
size_t pos, len, hostPos, hostLen;
if (firstDot != std::string::npos)
{
pos = firstDot + 1;
}
else
{
firstDot = pos = certName.size();
}
len = certName.size() - pos;
if (hostFirstDot != std::string::npos)
{
hostPos = hostFirstDot + 1;
}
else
{
hostFirstDot = hostPos = hostname.size();
}
hostLen = hostname.size() - hostPos;
// *. in the beginning of the cert name
if (certName.compare(0, firstDot, "*") == 0)
{
return certName.compare(pos, len, hostname, hostPos, hostLen) == 0;
}
// * in the left most. but other chars in the right
else if (certName[0] == '*')
{
// compare if `hostname` ends with `certName` but without the leftmost
// should be fine as domain names can't be that long
intmax_t hostnameIdx = hostname.size() - 1;
intmax_t certNameIdx = certName.size() - 1;
while (hostnameIdx >= 0 && certNameIdx != 0)
{
if (hostname[hostnameIdx] != certName[certNameIdx])
{
return false;
}
hostnameIdx--;
certNameIdx--;
}
if (certNameIdx != 0)
{
return false;
}
return true;
}
// * in the right of the first dot
else if (firstDot != 0 && certName[firstDot - 1] == '*')
{
if (certName.compare(pos, len, hostname, hostPos, hostLen) != 0)
{
return false;
}
for (size_t i = 0;
i < hostFirstDot && i < firstDot && certName[i] != '*';
i++)
{
if (hostname[i] != certName[i])
{
return false;
}
}
return true;
}
// else there's a * in the middle
else
{
if (certName.compare(pos, len, hostname, hostPos, hostLen) != 0)
{
return false;
}
for (size_t i = 0;
i < hostFirstDot && i < firstDot && certName[i] != '*';
i++)
{
if (hostname[i] != certName[i])
{
return false;
}
}
intmax_t hostnameIdx = hostFirstDot - 1;
intmax_t certNameIdx = firstDot - 1;
while (hostnameIdx >= 0 && certNameIdx >= 0 &&
certName[certNameIdx] != '*')
{
if (hostname[hostnameIdx] != certName[certNameIdx])
{
return false;
}
hostnameIdx--;
certNameIdx--;
}
return true;
}
assert(false && "This line should not be reached in verifySslName");
// should not reach
return certName == hostname;
}
#define STRINGIFY(x) #x #define STRINGIFY(x) #x
#define TOSTRING(x) STRINGIFY(x) #define TOSTRING(x) STRINGIFY(x)
@@ -171,15 +171,6 @@ inline std::string fromNativePath(const std::wstring &strPath)
return fromWidePath(strPath); return fromWidePath(strPath);
} }
/**
* @brief Check if the name supplied by the SSL Cert matches a FQDN
* @param certName The name supplied by the SSL Cert
* @param hostName The FQDN to match
*
* @return true if matches. false otherwise
*/
bool verifySslName(const std::string &certName, const std::string &hostName);
/** /**
* @brief Returns the TLS backend used by trantor. Could be "None", "OpenSSL" or * @brief Returns the TLS backend used by trantor. Could be "None", "OpenSSL" or
* "Botan" * "Botan"
+39 -4
View File
@@ -34,9 +34,36 @@ export DROGON_INSTALL="${THIRD_PARTY}/drogon/install/${TARGET_ARCH}"
# 由顶层 CMakeLists.txt 通过 CMAKE_PREFIX_PATH 注入查找路径。 # 由顶层 CMakeLists.txt 通过 CMAKE_PREFIX_PATH 注入查找路径。
DROGON_REPO="${THIRD_PARTY}/drogon_repo" DROGON_REPO="${THIRD_PARTY}/drogon_repo"
DROGON_CONFIG="${DROGON_INSTALL}/libs/cmake/Drogon/DrogonConfig.cmake" DROGON_CONFIG="${DROGON_INSTALL}/libs/cmake/Drogon/DrogonConfig.cmake"
EXPECTED_DROGON_VERSION="1.9.13"
EXPECTED_TRANTOR_VERSION="1.5.28"
if [ -f "${DROGON_CONFIG}" ]; then framework_versions_match() {
echo "[third_party] Drogon 已编译,跳过 (${DROGON_INSTALL})" local drogon_header="${DROGON_INSTALL}/include/drogon/version.h"
local drogon_cmake="${DROGON_INSTALL}/libs/cmake/Drogon/DrogonConfigVersion.cmake"
local trantor_cmake="${DROGON_INSTALL}/libs/cmake/Trantor/TrantorConfigVersion.cmake"
local drogon_link trantor_link
[ -f "${drogon_header}" ] &&
[ -f "${drogon_cmake}" ] &&
[ -f "${trantor_cmake}" ] &&
grep -Fq "#define DROGON_VERSION \"${EXPECTED_DROGON_VERSION}\"" "${drogon_header}" &&
grep -Fq "set(PACKAGE_VERSION \"${EXPECTED_DROGON_VERSION}\")" "${drogon_cmake}" &&
grep -Fq "set(PACKAGE_VERSION \"${EXPECTED_TRANTOR_VERSION}\")" "${trantor_cmake}" || return 1
drogon_link="$(readlink "${DROGON_INSTALL}/libs/libdrogon.so.1" 2>/dev/null || true)"
trantor_link="$(readlink "${DROGON_INSTALL}/libs/libtrantor.so.1" 2>/dev/null || true)"
[ "${drogon_link}" = "libdrogon.so.${EXPECTED_DROGON_VERSION}" ] &&
[ "${trantor_link}" = "libtrantor.so.${EXPECTED_TRANTOR_VERSION}" ]
}
if ! grep -Fq "set(DROGON_PATCH_VERSION 13)" "${DROGON_REPO}/CMakeLists.txt" ||
! grep -Fq "set(TRANTOR_PATCH_VERSION 28)" "${DROGON_REPO}/trantor/CMakeLists.txt"; then
echo "[third_party] 错误:Drogon 源码不是 ${EXPECTED_DROGON_VERSION} / Trantor ${EXPECTED_TRANTOR_VERSION}" >&2
return 1 2>/dev/null || exit 1
fi
if framework_versions_match; then
echo "[third_party] Drogon ${EXPECTED_DROGON_VERSION} / Trantor ${EXPECTED_TRANTOR_VERSION} 已编译,跳过 (${DROGON_INSTALL})"
else else
if [ ! -f "${DROGON_REPO}/CMakeLists.txt" ]; then if [ ! -f "${DROGON_REPO}/CMakeLists.txt" ]; then
echo "[third_party] ⚠️ 源码目录不存在: drogon_repo,跳过 Drogon" echo "[third_party] ⚠️ 源码目录不存在: drogon_repo,跳过 Drogon"
@@ -47,7 +74,11 @@ else
echo "[third_party] (后续 cmake find_package(Drogon) 将报错)" echo "[third_party] (后续 cmake find_package(Drogon) 将报错)"
return 0 2>/dev/null || exit 0 return 0 2>/dev/null || exit 0
else else
echo "[third_party] 编译 Drogon + Trantor (产物缺失) → ${DROGON_INSTALL}" if [ -e "${DROGON_INSTALL}" ]; then
echo "[third_party] 框架产物版本不匹配,重建 ${DROGON_INSTALL}"
rm -rf "${DROGON_INSTALL}"
fi
echo "[third_party] 编译 Drogon ${EXPECTED_DROGON_VERSION} + Trantor ${EXPECTED_TRANTOR_VERSION}${DROGON_INSTALL}"
DROGON_BUILD="${DROGON_REPO}/build" DROGON_BUILD="${DROGON_REPO}/build"
rm -rf "${DROGON_BUILD}" rm -rf "${DROGON_BUILD}"
@@ -70,7 +101,11 @@ else
cmake --install "${DROGON_BUILD}" cmake --install "${DROGON_BUILD}"
rm -rf "${DROGON_BUILD}" rm -rf "${DROGON_BUILD}"
echo "[third_party] ✅ Drogon + Trantor → ${DROGON_INSTALL}" if ! framework_versions_match; then
echo "[third_party] 错误:框架安装版本校验失败" >&2
return 1 2>/dev/null || exit 1
fi
echo "[third_party] ✅ Drogon ${EXPECTED_DROGON_VERSION} / Trantor ${EXPECTED_TRANTOR_VERSION}${DROGON_INSTALL}"
fi fi
fi fi