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
41 changes: 29 additions & 12 deletions src/pyrelation.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
#include "duckdb_python/map.hpp"
#include "duckdb_python/expression/pyexpression.hpp"
#include "duckdb/common/arrow/physical_arrow_collector.hpp"
#include "duckdb/common/file_system.hpp"
#include "duckdb_python/arrow/arrow_export_utils.hpp"

namespace duckdb {
Expand Down Expand Up @@ -1252,6 +1253,32 @@ static Value NestedDictToStruct(const nb::object &dictionary) {
return Value::STRUCT(std::move(children));
}

static void ApplyOverwriteOption(identifier_map_t<vector<Value>> &options, const string &function_name,
const string &filename, Relation &relation, const nb::object &overwrite) {
if (nb::none().is(overwrite)) {
return;
}
if (!nb::isinstance<nb::bool_>(overwrite)) {
throw InvalidInputException(function_name + " only accepts 'overwrite' as a boolean");
}
const bool overwrite_val = static_cast<bool>(nb::bool_(overwrite));
options["overwrite_or_ignore"] = {Value::BOOLEAN(overwrite_val)};
if (overwrite_val) {
return;
}
// Single-file COPY replaces via a tmp rename, so overwrite_or_ignore=false is a no-op.
auto context = relation.context->TryGetContext();
if (!context) {
throw InvalidInputException(function_name + " cannot run after the connection has been closed");
}
auto &fs = FileSystem::GetFileSystem(*context);
if (fs.FileExists(filename)) {
throw IOException("Cannot write to \"%s\" - it exists and is a file, not a directory! Enable OVERWRITE option "
"to overwrite the file",
filename);
}
}

void DuckDBPyRelation::ToParquet(const string &filename, const nb::object &compression, const nb::object &field_ids,
const nb::object &row_group_size_bytes, const nb::object &row_group_size,
const nb::object &overwrite, const nb::object &per_thread_output,
Expand Down Expand Up @@ -1327,12 +1354,7 @@ void DuckDBPyRelation::ToParquet(const string &filename, const nb::object &compr
options["append"] = {Value::BOOLEAN((bool)nb::bool_(append))};
}

if (!nb::none().is(overwrite)) {
if (!nb::isinstance<nb::bool_>(overwrite)) {
throw InvalidInputException("to_parquet only accepts 'overwrite' as a boolean");
}
options["overwrite_or_ignore"] = {Value::BOOLEAN((bool)nb::bool_(overwrite))};
}
ApplyOverwriteOption(options, "to_parquet", filename, *rel, overwrite);

if (!nb::none().is(per_thread_output)) {
if (!nb::isinstance<nb::bool_>(per_thread_output)) {
Expand Down Expand Up @@ -1467,12 +1489,7 @@ void DuckDBPyRelation::ToCSV(const string &filename, const nb::object &sep, cons
options["compression"] = {Value(nb::cast<std::string>(compression))};
}

if (!nb::none().is(overwrite)) {
if (!nb::isinstance<nb::bool_>(overwrite)) {
throw InvalidInputException("to_csv only accepts 'overwrite' as a boolean");
}
options["overwrite_or_ignore"] = {Value::BOOLEAN((bool)nb::bool_(overwrite))};
}
ApplyOverwriteOption(options, "to_csv", filename, *rel, overwrite);

if (!nb::none().is(per_thread_output)) {
if (!nb::isinstance<nb::bool_>(per_thread_output)) {
Expand Down
13 changes: 13 additions & 0 deletions tests/fast/api/test_to_csv.py
Original file line number Diff line number Diff line change
Expand Up @@ -264,6 +264,19 @@ def test_to_csv_overwrite_not_enabled(self):
with pytest.raises(duckdb.IOException, match="OVERWRITE"):
rel.to_csv(temp_file_name, header=True, partition_by=["c_category_1"])

def test_to_csv_overwrite_false_errors_on_existing_file(self):
temp_file_name = os.path.join(tempfile.mkdtemp(), next(tempfile._get_candidate_names())) # noqa: PTH118
duckdb.sql("SELECT 1 AS test").to_csv(temp_file_name, overwrite=False)
with pytest.raises(duckdb.IOException, match="OVERWRITE"):
duckdb.sql("SELECT 2 AS test").to_csv(temp_file_name, overwrite=False)
assert duckdb.read_csv(temp_file_name).fetchall() == [(1,)]

def test_to_csv_overwrite_true_replaces_existing_file(self):
temp_file_name = os.path.join(tempfile.mkdtemp(), next(tempfile._get_candidate_names())) # noqa: PTH118
duckdb.sql("SELECT 1 AS test").to_csv(temp_file_name)
duckdb.sql("SELECT 2 AS test").to_csv(temp_file_name, overwrite=True)
assert duckdb.read_csv(temp_file_name).fetchall() == [(2,)]

def test_to_csv_per_thread_output(self):
temp_file_name = os.path.join(tempfile.mkdtemp(), next(tempfile._get_candidate_names())) # noqa: PTH118
num_threads = duckdb.sql("select current_setting('threads')").fetchone()[0]
Expand Down
13 changes: 13 additions & 0 deletions tests/fast/api/test_to_parquet.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,19 @@ def test_overwrite(self, write_columns):

assert result.execute().fetchall() == expected

def test_overwrite_false_errors_on_existing_file(self):
temp_file_name = os.path.join(tempfile.mkdtemp(), next(tempfile._get_candidate_names())) # noqa: PTH118
duckdb.sql("SELECT 1 AS test").to_parquet(temp_file_name, overwrite=False)
with pytest.raises(duckdb.IOException, match="OVERWRITE"):
duckdb.sql("SELECT 2 AS test").to_parquet(temp_file_name, overwrite=False)
assert duckdb.read_parquet(temp_file_name).fetchall() == [(1,)]

def test_overwrite_true_replaces_existing_file(self):
temp_file_name = os.path.join(tempfile.mkdtemp(), next(tempfile._get_candidate_names())) # noqa: PTH118
duckdb.sql("SELECT 1 AS test").to_parquet(temp_file_name)
duckdb.sql("SELECT 2 AS test").to_parquet(temp_file_name, overwrite=True)
assert duckdb.read_parquet(temp_file_name).fetchall() == [(2,)]

def test_use_tmp_file(self):
temp_file_name = os.path.join(tempfile.mkdtemp(), next(tempfile._get_candidate_names())) # noqa: PTH118
df = pd.DataFrame(
Expand Down