Browse Source

fix(str): use character indices in find, index, ljust and rjust (#547)

find and index returned UTF-8 byte offsets and ljust/rjust measured width
in bytes, so on non-ASCII text s[s.find(t):] sliced at the wrong place and
'é'.ljust(3) padded one space short.

Co-authored-by: youdie006 <youdie006@users.noreply.github.com>
KBS 4 hours ago
parent
commit
554adf2ca6
2 changed files with 35 additions and 14 deletions
  1. 21 14
      src/bindings/py_str.c
  2. 14 0
      tests/041_str.py

+ 21 - 14
src/bindings/py_str.c

@@ -464,20 +464,20 @@ static bool str_zfill(int argc, py_Ref argv) {
 
 
 static bool str__widthjust_impl(bool left, int argc, py_Ref argv) {
 static bool str__widthjust_impl(bool left, int argc, py_Ref argv) {
     if(argc > 1 + 2) return TypeError("expected at most 2 arguments");
     if(argc > 1 + 2) return TypeError("expected at most 2 arguments");
-    char pad;
+    c11_sv pad;
     if(argc == 1 + 1) {
     if(argc == 1 + 1) {
-        pad = ' ';
+        pad = (c11_sv){" ", 1};
     } else {
     } else {
         if(!py_checkstr(&argv[2])) return false;
         if(!py_checkstr(&argv[2])) return false;
-        c11_string* padstr = pk_tostr(&argv[2]);
-        if(padstr->size != 1)
+        pad = c11_string__sv(pk_tostr(&argv[2]));
+        if(c11_sv__u8_length(pad) != 1)
             return TypeError("The fill character must be exactly one character long");
             return TypeError("The fill character must be exactly one character long");
-        pad = padstr->data[0];
     }
     }
     c11_sv self = c11_string__sv(pk_tostr(&argv[0]));
     c11_sv self = c11_string__sv(pk_tostr(&argv[0]));
     PY_CHECK_ARG_TYPE(1, tp_int);
     PY_CHECK_ARG_TYPE(1, tp_int);
     int width = py_toint(py_arg(1));
     int width = py_toint(py_arg(1));
-    if(width <= self.size) {
+    int delta = width - c11_sv__u8_length(self);
+    if(delta <= 0) {
         *py_retval() = argv[0];
         *py_retval() = argv[0];
         return true;
         return true;
     }
     }
@@ -485,12 +485,12 @@ static bool str__widthjust_impl(bool left, int argc, py_Ref argv) {
     c11_sbuf__ctor(&buf);
     c11_sbuf__ctor(&buf);
     if(left) {
     if(left) {
         c11_sbuf__write_sv(&buf, self);
         c11_sbuf__write_sv(&buf, self);
-        for(int i = 0; i < width - self.size; i++) {
-            c11_sbuf__write_char(&buf, pad);
+        for(int i = 0; i < delta; i++) {
+            c11_sbuf__write_sv(&buf, pad);
         }
         }
     } else {
     } else {
-        for(int i = 0; i < width - self.size; i++) {
-            c11_sbuf__write_char(&buf, pad);
+        for(int i = 0; i < delta; i++) {
+            c11_sbuf__write_sv(&buf, pad);
         }
         }
         c11_sbuf__write_sv(&buf, self);
         c11_sbuf__write_sv(&buf, self);
     }
     }
@@ -504,17 +504,24 @@ static bool str_rjust(int argc, py_Ref argv) { return str__widthjust_impl(false,
 
 
 static bool str_find(int argc, py_Ref argv) {
 static bool str_find(int argc, py_Ref argv) {
     if(argc > 3) return TypeError("find() takes at most 3 arguments");
     if(argc > 3) return TypeError("find() takes at most 3 arguments");
-    c11_string* self = pk_tostr(&argv[0]);
+    c11_sv self = c11_string__sv(pk_tostr(&argv[0]));
+    int u8_len = c11_sv__u8_length(self);
     int start = 0;
     int start = 0;
     if(argc == 3) {
     if(argc == 3) {
         PY_CHECK_ARG_TYPE(2, tp_int);
         PY_CHECK_ARG_TYPE(2, tp_int);
         start = py_toint(py_arg(2));
         start = py_toint(py_arg(2));
-        if(start < 0) start += c11_sv__u8_length(c11_string__sv(self));
+        if(start < 0) start += u8_len;
         if(start < 0) start = 0;
         if(start < 0) start = 0;
     }
     }
     PY_CHECK_ARG_TYPE(1, tp_str);
     PY_CHECK_ARG_TYPE(1, tp_str);
-    c11_string* sub = pk_tostr(&argv[1]);
-    int res = c11_sv__index2(c11_string__sv(self), c11_string__sv(sub), start);
+    c11_sv sub = c11_string__sv(pk_tostr(&argv[1]));
+    if(start > u8_len) {
+        py_newint(py_retval(), -1);
+        return true;
+    }
+    // start and the result are character indices, not byte offsets
+    int res = c11_sv__index2(self, sub, c11__unicode_index_to_byte(self.data, start));
+    if(res != -1) res = c11__byte_index_to_unicode(self.data, res);
     py_newint(py_retval(), res);
     py_newint(py_retval(), res);
     return true;
     return true;
 }
 }

+ 14 - 0
tests/041_str.py

@@ -157,6 +157,10 @@ assert a.rjust(5) == '  123'
 assert a.rjust(5, '0') == '00123'
 assert a.rjust(5, '0') == '00123'
 assert a.ljust(5) == '123  '
 assert a.ljust(5) == '123  '
 assert a.ljust(5, '0') == '12300'
 assert a.ljust(5, '0') == '12300'
+assert 'é'.ljust(3) == 'é  '
+assert 'é'.rjust(3, '*') == '**é'
+assert '日本'.ljust(3, '-') == '日本-'
+assert 'a'.rjust(3, 'é') == 'ééa'
 
 
 assert '\x30\x31\x32' == '012'
 assert '\x30\x31\x32' == '012'
 assert '\b\b\b' == '\x08\x08\x08'
 assert '\b\b\b' == '\x08\x08\x08'
@@ -181,6 +185,16 @@ assert a.index('2', -2) == 1
 assert a.find('1') == 0
 assert a.find('1') == 0
 assert a.find('1', 1) == -1
 assert a.find('1', 1) == -1
 
 
+# find and index return character indices, not byte offsets
+s = 'aßa-é日'
+assert s.find('é') == 4
+assert s.find('日', -1) == 5
+assert s.find('a', 1) == 2
+assert s.index('é', 3) == 4
+assert s[s.find('é'):] == 'é日'
+assert 'ab'.find('', 2) == 2
+assert 'ab'.find('', 3) == -1
+
 b = list("Hello, World!")
 b = list("Hello, World!")
 assert b == ['H', 'e', 'l', 'l', 'o', ',', ' ', 'W', 'o', 'r', 'l', 'd', '!']
 assert b == ['H', 'e', 'l', 'l', 'o', ',', ' ', 'W', 'o', 'r', 'l', 'd', '!']
 assert b[::-1] == ['!', 'd', 'l', 'r', 'o', 'W', ' ', ',', 'o', 'l', 'l', 'e', 'H']
 assert b[::-1] == ['!', 'd', 'l', 'r', 'o', 'W', ' ', ',', 'o', 'l', 'l', 'e', 'H']