Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 41 additions & 32 deletions src-tauri/src/drivers/mysql/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -746,26 +746,12 @@ pub async fn delete_record(
.await
}

pub async fn update_record(
params: &ConnectionParams,
table: &str,
pk_map: &HashMap<String, serde_json::Value>,
col_name: &str,
new_val: serde_json::Value,
fn push_mysql_update_value(
qb: &mut sqlx::QueryBuilder<'_, sqlx::MySql>,
new_val: &serde_json::Value,
text: TextProto,
max_blob_size: u64,
) -> Result<u64, String> {
let pool = get_mysql_pool(params).await?;
// Behind a prepared-statement-less bastion every value is inlined as an
// escaped literal instead of bound (see `force_text_protocol`).
let text = resolve_text_proto(&pool, params).await?;
let pk_pairs = build_mysql_pk_where(pk_map)?;

let mut qb = sqlx::QueryBuilder::new(format!(
"UPDATE `{}` SET `{}` = ",
escape_identifier(table),
escape_identifier(col_name)
));

) -> Result<(), String> {
match new_val {
serde_json::Value::Number(n) => {
if n.is_i64() {
Expand All @@ -784,7 +770,7 @@ pub async fn update_record(
if s == "__USE_DEFAULT__" {
qb.push("DEFAULT");
} else if let Some(bytes) =
crate::drivers::common::decode_blob_wire_format(&s, max_blob_size)
crate::drivers::common::decode_blob_wire_format(s, max_blob_size)
{
// Blob wire format: decode to raw bytes so the DB stores binary data,
// not the internal wire format string.
Expand All @@ -793,17 +779,17 @@ pub async fn update_record(
} else {
qb.push_bind(bytes);
}
} else if is_raw_sql_function(&s) {
} else if is_raw_sql_function(s) {
qb.push(s);
} else if is_wkt_geometry(&s) {
} else if is_wkt_geometry(s) {
qb.push("ST_GeomFromText(");
if text.enabled {
qb.push(mysql_string_literal(&s, text.no_backslash_escapes));
qb.push(mysql_string_literal(s, text.no_backslash_escapes));
} else {
qb.push_bind(s);
qb.push_bind(s.clone());
}
qb.push(")");
} else if let Some(n) = parse_unsafe_bigint_string(&s) {
} else if let Some(n) = parse_unsafe_bigint_string(s) {
// Bigints outside JS safe range come back from the UI as strings
// (see drivers::common::i64_to_json). Bind them as native i64 so
// BIGINT columns receive the exact value.
Expand All @@ -813,33 +799,56 @@ pub async fn update_record(
qb.push_bind(n);
}
} else if text.enabled {
qb.push(mysql_string_literal(&s, text.no_backslash_escapes));
qb.push(mysql_string_literal(s, text.no_backslash_escapes));
} else {
qb.push_bind(s);
qb.push_bind(s.clone());
}
}
serde_json::Value::Bool(b) => {
if text.enabled {
qb.push(if b { "1" } else { "0" });
qb.push(if *b { "1" } else { "0" });
} else {
qb.push_bind(b);
qb.push_bind(*b);
}
}
serde_json::Value::Null => {
qb.push("NULL");
}
serde_json::Value::Object(_) | serde_json::Value::Array(_) => {
let json_str = serde_json::to_string(&new_val).map_err(|e| e.to_string())?;
qb.push("CAST(");
let json_str = serde_json::to_string(new_val).map_err(|e| e.to_string())?;
if text.enabled {
qb.push(mysql_string_literal(&json_str, text.no_backslash_escapes));
} else {
qb.push_bind(json_str);
}
qb.push(" AS JSON)");
}
}

Ok(())
}

pub async fn update_record(
params: &ConnectionParams,
table: &str,
pk_map: &HashMap<String, serde_json::Value>,
col_name: &str,
new_val: serde_json::Value,
max_blob_size: u64,
) -> Result<u64, String> {
let pool = get_mysql_pool(params).await?;
// Behind a prepared-statement-less bastion every value is inlined as an
// escaped literal instead of bound (see `force_text_protocol`).
let text = resolve_text_proto(&pool, params).await?;
let pk_pairs = build_mysql_pk_where(pk_map)?;

let mut qb = sqlx::QueryBuilder::new(format!(
"UPDATE `{}` SET `{}` = ",
escape_identifier(table),
escape_identifier(col_name)
));

push_mysql_update_value(&mut qb, &new_val, text, max_blob_size)?;

qb.push(" WHERE ");
let mut first = true;
for (col, val) in &pk_pairs {
Expand Down
32 changes: 31 additions & 1 deletion src-tauri/src/drivers/mysql/tests.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
use super::build_mysql_pk_where;
use super::{is_text_protocol_stmt, MysqlDriver};
use super::{is_text_protocol_stmt, push_mysql_update_value, MysqlDriver, TextProto};
use super::helpers::{inline_str_placeholders, mysql_bytes_literal, mysql_string_literal};
use crate::drivers::driver_trait::DatabaseDriver;
use crate::models::{ConnectionParams, DatabaseSelection};
Expand Down Expand Up @@ -72,6 +72,36 @@ fn mysql_bytes_literal_hex_encodes() {
assert_eq!(mysql_bytes_literal(b"AB"), "x'4142'");
}

#[test]
fn mysql_json_update_value_binds_without_json_cast() {
let mut qb = sqlx::QueryBuilder::<sqlx::MySql>::new("SET `payload` = ");

push_mysql_update_value(
&mut qb,
&serde_json::json!({ "ok": true }),
TextProto::PREPARED,
1024,
)
.unwrap();

assert_eq!(qb.sql(), "SET `payload` = ?");
}

#[test]
fn mysql_json_update_value_inlines_without_json_cast_in_text_protocol() {
let mut qb = sqlx::QueryBuilder::<sqlx::MySql>::new("SET `payload` = ");

push_mysql_update_value(
&mut qb,
&serde_json::json!({ "ok": true }),
TextProto::protocol_only(true),
1024,
)
.unwrap();

assert_eq!(qb.sql(), "SET `payload` = '{\\\"ok\\\":true}'");
}

#[test]
fn inline_str_placeholders_substitutes_in_order() {
let sql = "WHERE table_schema = ? AND table_name = ?";
Expand Down
24 changes: 21 additions & 3 deletions src/components/modals/ErrorModal.tsx
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
import { useState } from "react";
import { useTranslation } from "react-i18next";
import { AlertTriangle, X } from "lucide-react";
import { AlertTriangle, Check, Copy, X } from "lucide-react";
import { Modal } from "../ui/Modal";
import { copyTextToClipboard } from "../../utils/clipboard";

interface ErrorModalProps {
isOpen: boolean;
Expand All @@ -10,6 +12,13 @@ interface ErrorModalProps {

export const ErrorModal = ({ isOpen, onClose, message }: ErrorModalProps) => {
const { t } = useTranslation();
const [copied, setCopied] = useState(false);

const handleCopy = async () => {
await copyTextToClipboard(message);
setCopied(true);
window.setTimeout(() => setCopied(false), 2000);
};

return (
<Modal isOpen={isOpen} onClose={onClose}>
Expand All @@ -32,10 +41,19 @@ export const ErrorModal = ({ isOpen, onClose, message }: ErrorModalProps) => {
</div>

<div className="p-6 overflow-y-auto">
<p className="text-sm text-secondary break-words">{message}</p>
<pre className="text-sm text-secondary whitespace-pre-wrap break-words select-text font-mono">
{message}
</pre>
</div>

<div className="p-4 border-t border-default bg-base/50 flex justify-end">
<div className="p-4 border-t border-default bg-base/50 flex justify-end gap-3">
<button
onClick={handleCopy}
className="flex items-center gap-2 px-4 py-2 text-secondary hover:text-primary border border-strong rounded-lg text-sm font-medium transition-colors"
>
{copied ? <Check size={15} className="text-green-400" /> : <Copy size={15} />}
{copied ? t("dataGrid.copied") : t("common.copy")}
</button>
<button
onClick={onClose}
className="px-4 py-2 bg-blue-600 hover:bg-blue-500 text-white rounded-lg text-sm font-medium transition-colors"
Expand Down
33 changes: 11 additions & 22 deletions src/components/modals/QuickNavigatorModal.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@ import { invoke } from "@tauri-apps/api/core";
import { useDatabase } from "../../hooks/useDatabase";
import { useAlert } from "../../hooks/useAlert";
import { quoteTableRef } from "../../utils/identifiers";
import { isMultiDatabaseCapable, getDatabaseList } from "../../utils/database";
import { isMultiDatabaseCapable, usesMultiDatabaseLayout } from "../../utils/database";
import { getNavigatorItems, filterNavigatorItems } from "../../utils/quickNavigator";
import { newConsoleForTable } from "../../utils/newConsole";
import type { RoutineInfo, TriggerInfo } from "../../contexts/DatabaseContext";
Expand Down Expand Up @@ -39,54 +39,43 @@ export const QuickNavigatorModal = ({ isOpen, onClose, onGenerateSql, onInspect
schemas,
loadSchemaData,
loadDatabaseData,
connections,
selectedDatabases,
} = useDatabase();

const [search, setSearch] = useState("");
const [selectedIndex, setSelectedIndex] = useState(0);

// Find the active connection configuration
const activeConn = useMemo(() => {
return connections.find((c) => c.id === activeConnectionId);
}, [connections, activeConnectionId]);

// Resolve the configured database list for the connection
const configuredDatabases = useMemo(() => {
if (!activeConn) return [];
return getDatabaseList(activeConn.params.database);
}, [activeConn]);

// Load metadata for all schemas/databases in the background when the modal is open
useEffect(() => {
if (!isOpen || !activeConnectionId) return;

const loadAll = async () => {
const isMultiDb = isMultiDatabaseCapable(activeCapabilities);
const isMultiDb = usesMultiDatabaseLayout(activeCapabilities, selectedDatabases);
const hasSchemas = activeCapabilities?.schemas;

if (hasSchemas && schemas) {
schemas.forEach((schema) => {
loadSchemaData(schema);
});
} else if (isMultiDb && configuredDatabases) {
configuredDatabases.forEach((db) => {
} else if (isMultiDb) {
selectedDatabases.forEach((db) => {
loadDatabaseData(db);
});
}
};

loadAll();
}, [isOpen, activeConnectionId, activeCapabilities, schemas, configuredDatabases, loadSchemaData, loadDatabaseData]);
}, [isOpen, activeConnectionId, activeCapabilities, schemas, selectedDatabases, loadSchemaData, loadDatabaseData]);

// Gather all schema items based on database capabilities
const items = useMemo(() => {
return getNavigatorItems({
activeConnectionId,
hasSchemas: !!activeCapabilities?.schemas,
isMultiDb: isMultiDatabaseCapable(activeCapabilities),
isMultiDb: usesMultiDatabaseLayout(activeCapabilities, selectedDatabases),
schemas,
schemaDataMap,
configuredDatabases,
configuredDatabases: selectedDatabases,
databaseDataMap,
tables,
views,
Expand All @@ -97,9 +86,9 @@ export const QuickNavigatorModal = ({ isOpen, onClose, onGenerateSql, onInspect
}, [
activeConnectionId,
activeCapabilities,
selectedDatabases,
schemas,
schemaDataMap,
configuredDatabases,
databaseDataMap,
tables,
views,
Expand All @@ -110,10 +99,10 @@ export const QuickNavigatorModal = ({ isOpen, onClose, onGenerateSql, onInspect

// Check if we have multiple databases/schemas to show group headers
const showGroupHeaders = useMemo(() => {
const isMultiDb = isMultiDatabaseCapable(activeCapabilities);
const isMultiDb = usesMultiDatabaseLayout(activeCapabilities, selectedDatabases);
const hasSchemas = activeCapabilities?.schemas;
return !!(hasSchemas || isMultiDb);
}, [activeCapabilities]);
}, [activeCapabilities, selectedDatabases]);

// Dynamically filter items as user types
const filteredItems = useMemo(() => {
Expand Down
30 changes: 30 additions & 0 deletions tests/components/modals/ErrorModal.test.tsx
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import { ErrorModal } from "../../../src/components/modals/ErrorModal";

describe("ErrorModal", () => {
it("renders selectable error text and copies it", async () => {
const writeText = vi.fn().mockResolvedValue(undefined);
Object.defineProperty(navigator, "clipboard", {
value: { writeText },
configurable: true,
});

render(
<ErrorModal
isOpen
onClose={vi.fn()}
message={"Query failed\nUnknown column payment_status"}
/>,
);

const errorText = screen.getByText(/Unknown column payment_status/);
expect(errorText).toHaveClass("select-text");

fireEvent.click(screen.getByRole("button", { name: /common.copy/i }));

await waitFor(() => {
expect(writeText).toHaveBeenCalledWith("Query failed\nUnknown column payment_status");
});
});
});